diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..9eca067 --- /dev/null +++ b/.env.example @@ -0,0 +1,24 @@ +# A-Share Data Platform Configuration + +# Data directory (absolute or relative to project root) +DATA_DIR=data + +# DuckDB database path +DUCKDB_PATH=data/duckdb/ashare.db + +# API server +API_HOST=0.0.0.0 +API_PORT=8000 + +# Backfill concurrency (ThreadPoolExecutor max_workers) +BACKFILL_WORKERS=10 + +# AKShare retry settings +AKSHARE_MAX_RETRIES=3 +AKSHARE_RETRY_DELAY=1.0 + +# Real-time polling interval (seconds) +REALTIME_POLL_INTERVAL=5 + +# Log level: DEBUG, INFO, WARNING, ERROR +LOG_LEVEL=INFO diff --git a/.gitignore b/.gitignore index 83972fa..543d9c7 100644 --- a/.gitignore +++ b/.gitignore @@ -1,218 +1,35 @@ -# Byte-compiled / optimized / DLL files +# Python __pycache__/ -*.py[codz] -*$py.class - -# C extensions -*.so - -# Distribution / packaging -.Python -build/ -develop-eggs/ -dist/ -downloads/ -eggs/ -.eggs/ -lib/ -lib64/ -parts/ -sdist/ -var/ -wheels/ -share/python-wheels/ +*.py[cod] *.egg-info/ -.installed.cfg -*.egg -MANIFEST +dist/ +build/ -# PyInstaller -# Usually these files are written by a python script from a template -# before PyInstaller builds the exe, so as to inject date/other infos into it. -*.manifest -*.spec - -# Installer logs -pip-log.txt -pip-delete-this-directory.txt - -# Unit test / coverage reports -htmlcov/ -.tox/ -.nox/ -.coverage -.coverage.* -.cache -nosetests.xml -coverage.xml -*.cover -*.py.cover -.hypothesis/ -.pytest_cache/ -cover/ - -# Translations -*.mo -*.pot - -# Django stuff: -*.log -local_settings.py -db.sqlite3 -db.sqlite3-journal - -# Flask stuff: -instance/ -.webassets-cache - -# Scrapy stuff: -.scrapy - -# Sphinx documentation -docs/_build/ - -# PyBuilder -.pybuilder/ -target/ - -# Jupyter Notebook -.ipynb_checkpoints - -# IPython -profile_default/ -ipython_config.py - -# pyenv -# For a library or package, you might want to ignore these files since the code is -# intended to run in multiple environments; otherwise, check them in: -# .python-version - -# pipenv -# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control. -# However, in case of collaboration, if having platform-specific dependencies or dependencies -# having no cross-platform support, pipenv may install dependencies that don't work, or not -# install all needed dependencies. -# Pipfile.lock - -# UV -# Similar to Pipfile.lock, it is generally recommended to include uv.lock in version control. -# This is especially recommended for binary packages to ensure reproducibility, and is more -# commonly ignored for libraries. -# uv.lock - -# poetry -# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control. -# This is especially recommended for binary packages to ensure reproducibility, and is more -# commonly ignored for libraries. -# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control -# poetry.lock -# poetry.toml - -# pdm -# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control. -# pdm recommends including project-wide configuration in pdm.toml, but excluding .pdm-python. -# https://pdm-project.org/en/latest/usage/project/#working-with-version-control -# pdm.lock -# pdm.toml -.pdm-python -.pdm-build/ - -# pixi -# Similar to Pipfile.lock, it is generally recommended to include pixi.lock in version control. -# pixi.lock -# Pixi creates a virtual environment in the .pixi directory, just like venv module creates one -# in the .venv directory. It is recommended not to include this directory in version control. -.pixi - -# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm -__pypackages__/ - -# Celery stuff -celerybeat-schedule -celerybeat.pid - -# Redis -*.rdb -*.aof -*.pid - -# RabbitMQ -mnesia/ -rabbitmq/ -rabbitmq-data/ - -# ActiveMQ -activemq-data/ - -# SageMath parsed files -*.sage.py - -# Environments +# Environment .env -.envrc -.venv -env/ +*.env.local + +# Data (large files) +data/parquet/ +data/duckdb/ + +# Claude +.claude/ + +# IDE +.idea/ +.vscode/ +*.swp +*.swo + +# OS +.DS_Store +Thumbs.db + +# Logs +*.log +logs/ + +# Virtual env +.venv/ venv/ -ENV/ -env.bak/ -venv.bak/ - -# Spyder project settings -.spyderproject -.spyproject - -# Rope project settings -.ropeproject - -# mkdocs documentation -/site - -# mypy -.mypy_cache/ -.dmypy.json -dmypy.json - -# Pyre type checker -.pyre/ - -# pytype static type analyzer -.pytype/ - -# Cython debug symbols -cython_debug/ - -# PyCharm -# JetBrains specific template is maintained in a separate JetBrains.gitignore that can -# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore -# and can be added to the global gitignore or merged into this file. For a more nuclear -# option (not recommended) you can uncomment the following to ignore the entire idea folder. -# .idea/ - -# Abstra -# Abstra is an AI-powered process automation framework. -# Ignore directories containing user credentials, local state, and settings. -# Learn more at https://abstra.io/docs -.abstra/ - -# Visual Studio Code -# Visual Studio Code specific template is maintained in a separate VisualStudioCode.gitignore -# that can be found at https://github.com/github/gitignore/blob/main/Global/VisualStudioCode.gitignore -# and can be added to the global gitignore or merged into this file. However, if you prefer, -# you could uncomment the following to ignore the entire vscode folder -# .vscode/ -# Temporary file for partial code execution -tempCodeRunnerFile.py - -# Ruff stuff: -.ruff_cache/ - -# PyPI configuration file -.pypirc - -# Marimo -marimo/_static/ -marimo/_lsp/ -__marimo__/ - -# Streamlit -.streamlit/secrets.toml diff --git a/CLAUDE.md b/CLAUDE.md new file mode 100644 index 0000000..b8e6e25 --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1,94 @@ +# CLAUDE.md - A-Share Data Platform + +## Project Overview + +A-share (Chinese stock market) data platform providing 9 K-line frequencies via Parquet + DuckDB storage with REST and WebSocket APIs. Greenfield project, currently at v0.1.0. + +## Tech Stack + +- **Data source**: `akshare` (wraps East Money and Sina APIs) +- **Storage**: Parquet (Zstd compression, Hive-partitioned `year=YYYY/month=MM/day=DD/`) +- **Query engine**: DuckDB (embedded OLAP, `read_parquet` with `hive_partitioning=true, union_by_name=true`) +- **Web**: FastAPI + uvicorn +- **CLI**: Typer +- **Scheduler**: APScheduler 3.x (AsyncIOScheduler) +- **Config**: pydantic-settings (reads `.env`) +- **Logging**: loguru + +## Project Structure + +``` +src/ashare_dp/ +├── config.py # Pydantic Settings (all config via env vars) +├── core/ +│ ├── models.py # Freq enum with storage_dir property, freq groupings +│ ├── calendar.py # Trading calendar, market state, Beijing TZ +│ └── exceptions.py # AShareDPError hierarchy +├── data/ +│ ├── akshare_client.py # AKShare wrapper: dual backend, retry, rate limit +│ ├── backfill.py # Historical backfill (daily/weekly/monthly + minute) +│ ├── eod.py # End-of-day batch pull (all stocks, all freqs) +│ └── realtime.py # Background async spot poller → WebSocket broadcast +├── storage/ +│ ├── database.py # DuckDB singleton (get_db) +│ ├── schema.py # DDL: stock_info, trading_calendar +│ ├── repository.py # KLineRepository: write_klines, read_klines, get_latest +│ └── partitioning.py # Hive partition path builder +├── api/ +│ ├── app.py # FastAPI factory + lifespan + embedded docs HTML +│ ├── deps.py # FastAPI DI (get_repo) +│ ├── routers/ # stocks, kline, realtime, calendar routers +│ └── websocket/ +│ ├── manager.py # ConnectionManager: subscribe/broadcast with async lock +│ └── handlers.py # WS message dispatch (/ws/realtime) +├── scheduler/ +│ ├── scheduler.py # APScheduler setup (EOD 15:05, health check 08:00) +│ └── jobs.py # Job implementations +└── cli/ + ├── main.py # Typer root: ashare-dp {backfill, serve, query, version} + ├── backfill_cmd.py + ├── serve_cmd.py + └── query_cmd.py +``` + +## Critical Design Decisions + +### Freq.storage_dir property +macOS APFS is case-insensitive, so `kline_1m` and `kline_1M` collide. The `Freq.M1` (monthly) uses `storage_dir = "1mon"` to disambiguate. Always use `freq.storage_dir` for filesystem paths, never `freq.value`. + +### Batched writes to avoid file overwrites +`write_klines` groups all stocks for a day into a single `data.parquet` file. The backfill and EOD pipelines collect all DataFrames first, then write once per day/freq. Never write per-stock-per-day files. + +### Dual backend with auto-fallback +`AKShareClient._resolve_backend()` probes East Money reachability once and caches the result. If unreachable (geo-blocked outside China), falls back to Sina. `get_stock_list()` and `get_trading_calendar()` are Sina-only. + +### SQL parameterization +`read_klines()` and `get_latest()` use DuckDB parameterized queries (`$1`, `$2`) for user-supplied values (ts_code, dates). Do NOT use f-string interpolation for user input. + +### 2h derivation +2-hour K-lines are derived on-the-fly from 1h data via `_read_2h()` using pandas resampling. No Parquet storage for 2h. + +## Key Conventions + +- Stock codes: `ts_code` format is `"000001.SZ"` (6-digit code + exchange suffix). Internal API calls use 6-digit numeric strings. +- Exchange mapping: codes starting with `6`/`9` → SH, `4`/`8`/`92` → BJ, rest → SZ +- Backend-specific column normalization: Sina returns English columns (`date`, `open`, etc.), EM returns Chinese. `_normalize_hist_df` and `_normalize_min_df` handle both. +- Proxy: env vars cleared + `urllib.request.getproxies` monkey-patched at module import time in `akshare_client.py` +- All async state in `ConnectionManager` is protected by `asyncio.Lock` + +## Running + +```bash +pip install -e ".[dev]" +ashare-dp backfill init # First time: create tables, load stock list +ashare-dp backfill daily # Backfill all daily/weekly/monthly history +ashare-dp serve start # Start API + scheduler + realtime poller +ashare-dp query stats # Check data status +``` + +## Tests + +```bash +pytest +ruff check src/ +``` diff --git a/README.md b/README.md new file mode 100644 index 0000000..70f91a1 --- /dev/null +++ b/README.md @@ -0,0 +1,186 @@ +# A-Share Data Platform + +A股全量数据服务 — Parquet + DuckDB 存储,REST + WebSocket 双协议。 + +## 支持的 K 线周期 + +| 周期 | 代码 | 数据范围 | 来源 | +|------|------|----------|------| +| 1 分钟 | `1m` | 近 1-3 月(需每日盘后积累) | AKShare | +| 5 分钟 | `5m` | 同上 | AKShare | +| 15 分钟 | `15m` | 同上 | AKShare | +| 30 分钟 | `30m` | 同上 | AKShare | +| 1 小时 | `1h` | 同上 | AKShare | +| 2 小时 | `2h` | 从 `1h` 实时推导 | DuckDB | +| 日线 | `1d` | 全部历史(1990 年起) | AKShare | +| 周线 | `1w` | 从日线推导 | DuckDB | +| 月线 | `1M` | 从日线推导 | DuckDB | + +分钟线数据受限于上游 API 只保留近 1-3 个月,必须通过**每日盘后自动拉取**持续积累。日线/周线/月线可随时回填全部历史。 + +## 架构 + +``` +AKShare (East Money / Sina) + → AKShareClient (dual backend, auto-fallback) + → BackfillPipeline / EODPipeline + → Parquet (Hive-partitioned, Zstd compressed) + → DuckDB (metadata + read_parquet queries) + → FastAPI (REST + WebSocket) +``` + +- **存储**: Parquet 列存(Zstd 压缩 ~80%),Hive 分区 `year=YYYY/month=MM/day=DD/data.parquet` +- **查询**: DuckDB 内嵌 OLAP,`read_parquet()` 直接读取,支持分区裁剪和谓词下推 +- **数据源**: AKShare 封装,双后端自动切换(East Money 国内优先,Sina 全球可访问) + +## 快速开始 + +### 安装 + +```bash +pip install -e . +``` + +### 配置 + +```bash +cp .env.example .env +# 编辑 .env 按需调整参数 +``` + +### 初始化数据库 + 导入股票列表 + +```bash +ashare-dp backfill init +``` + +### 回填历史数据 + +```bash +# 回填全部日线/周线/月线 +ashare-dp backfill daily --workers 10 + +# 回填指定股票指定日期范围 +ashare-dp backfill daily --start 2025-01-01 --end 2026-05-16 --symbols 000001,600000 + +# 回填近 30 天分钟数据 +ashare-dp backfill minute --days 30 --workers 5 +``` + +### 启动 API 服务 + +```bash +ashare-dp serve start --port 8000 +``` + +启动后访问 `http://localhost:8000/` 查看文档,`http://localhost:8000/docs` 查看 Swagger。 + +### 命令行查询 + +```bash +ashare-dp query kline 1d 000001.SZ --start 2026-01-01 +ashare-dp query latest --freq 1d --ts-code 000001.SZ +ashare-dp query stats +ashare-dp query stocks --exchange SH +``` + +## API 端点 + +所有 REST 端点前缀 `/api/v1`。 + +### 股票查询 + +| 方法 | 路径 | 说明 | +|------|------|------| +| GET | `/stocks` | 分页列表,可按交易所/板块筛选 | +| GET | `/stocks/search?q=平安` | 名称/代码模糊搜索 | +| GET | `/stocks/{ts_code}` | 单只股票详情 | + +### K 线查询 + +| 方法 | 路径 | 说明 | +|------|------|------| +| GET | `/klines/{freq}?ts_code=&start_date=&end_date=` | 单只 K 线查询 | +| POST | `/klines/{freq}/batch` | 批量 K 线查询 | +| GET | `/klines/{freq}/latest?ts_code=` | 最新交易日数据 | +| GET | `/klines/available-freqs` | 支持的频率列表 | + +### 实时行情 + +| 方法 | 路径 | 说明 | +|------|------|------| +| GET | `/realtime/spot?codes=000001.SZ,600000.SH` | 实时快照 | +| GET | `/realtime/market-state` | 市场状态 | + +### 交易日历 + +| 方法 | 路径 | 说明 | +|------|------|------| +| GET | `/calendar/trading-days?start=&end=` | 区间内交易日 | +| GET | `/calendar/is-trading-day?date=` | 判断交易日 | +| GET | `/calendar/next-trading-day?date=` | 下一个交易日 | + +### WebSocket + +``` +ws://localhost:8000/ws/realtime +``` + +交易时段每 5 秒推送订阅股票的实时行情。连接后发送 JSON 控制消息: + +```json +{"action": "subscribe", "codes": ["000001.SZ", "600519.SH"]} +{"action": "unsubscribe", "codes": ["000001.SZ"]} +{"action": "unsubscribe_all"} +``` + +### 其他 + +| 方法 | 路径 | 说明 | +|------|------|------| +| GET | `/health` | 健康检查 | +| GET | `/stats` | 数据库统计(各周期记录数、日期范围) | + +## CLI 命令 + +``` +ashare-dp version # 显示版本 +ashare-dp backfill init # 初始化数据库 +ashare-dp backfill daily [...] # 回填日线/周线/月线 +ashare-dp backfill minute [...] # 回填分钟线 +ashare-dp serve start [...] # 启动 API 服务 +ashare-dp query kline ... # 查询 K 线 +ashare-dp query latest ... # 最新数据 +ashare-dp query stocks ... # 股票列表 +ashare-dp query stats # 数据统计 +``` + +## 配置参数 + +| 变量 | 默认值 | 说明 | +|------|--------|------| +| `DATA_DIR` | `data` | 数据目录 | +| `DUCKDB_PATH` | `data/duckdb/ashare.db` | DuckDB 文件路径 | +| `API_HOST` | `0.0.0.0` | API 绑定地址 | +| `API_PORT` | `8000` | API 绑定端口 | +| `BACKFILL_WORKERS` | `10` | 回填并发线程数 | +| `AKSHARE_MAX_RETRIES` | `3` | API 调用重试次数 | +| `AKSHARE_RETRY_DELAY` | `1.0` | 重试基础延迟(指数退避) | +| `AKSHARE_BACKEND` | `auto` | 数据后端: `auto`, `em` (East Money), `sina` | +| `REALTIME_POLL_INTERVAL` | `5` | 实时行情轮询间隔(秒) | +| `LOG_LEVEL` | `INFO` | 日志级别 | + +## 部署说明 + +- **国内服务器**: 配置 `AKSHARE_BACKEND=em` 使用 East Money 后端(数据质量更好) +- **海外服务器**: 保持 `auto`,客户端会自动检测并降级到 Sina 后端 +- **分钟线积累**: 盘后拉取任务(15:05 北京时间)必须稳定运行,否则分钟线历史会出现缺口 +- **存储估算**: 约 12 GB / ~3,300 文件(全部历史 + 所有频率),建议 SSD + +## 开发 + +```bash +pip install -e ".[dev]" +pytest +ruff check src/ +``` diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..a38bc4c --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,38 @@ +[project] +name = "ashare-dp" +version = "0.1.0" +description = "A-Share Data Platform: Parquet + DuckDB storage with REST/WebSocket APIs" +requires-python = ">=3.11" +dependencies = [ + "akshare>=1.17.0", + "duckdb>=1.2.0", + "fastapi>=0.115.0", + "uvicorn[standard]>=0.34.0", + "pydantic>=2.0", + "pydantic-settings>=2.0", + "apscheduler>=3.10.0", + "typer>=0.15.0", + "loguru>=0.7.0", + "httpx>=0.28.0", + "pyarrow>=18.0.0", + "pandas>=2.0", + "pytz>=2024.0", +] + +[project.optional-dependencies] +dev = [ + "pytest>=8.0", + "pytest-asyncio>=0.25.0", + "pytest-cov>=6.0", + "ruff>=0.9.0", +] + +[project.scripts] +ashare-dp = "ashare_dp.cli.main:app" + +[build-system] +requires = ["setuptools>=75.0"] +build-backend = "setuptools.build_meta" + +[tool.setuptools.packages.find] +where = ["src"] diff --git a/src/ashare_dp/__init__.py b/src/ashare_dp/__init__.py new file mode 100644 index 0000000..8221b90 --- /dev/null +++ b/src/ashare_dp/__init__.py @@ -0,0 +1,3 @@ +"""A-Share Data Platform: Parquet + DuckDB storage with REST/WebSocket APIs.""" + +__version__ = "0.1.0" diff --git a/src/ashare_dp/api/__init__.py b/src/ashare_dp/api/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/ashare_dp/api/app.py b/src/ashare_dp/api/app.py new file mode 100644 index 0000000..aec3452 --- /dev/null +++ b/src/ashare_dp/api/app.py @@ -0,0 +1,388 @@ +"""FastAPI application factory with lifespan management.""" + +from __future__ import annotations + +from contextlib import asynccontextmanager + +from fastapi import FastAPI +from fastapi.responses import HTMLResponse +from loguru import logger + +from ashare_dp.api.routers import stocks, kline, realtime, calendar +from ashare_dp.core.models import Freq +from ashare_dp.storage.repository import KLineRepository +from ashare_dp.api.websocket.handlers import router as ws_router +from ashare_dp.storage.database import get_db +from ashare_dp.storage.schema import DDL_STATEMENTS + +DOCS_HTML = r""" + + + + + +A-Share Data Platform + + + +
+ +

A-Share Data Platform

+

A股全量数据服务 — Parquet + DuckDB 存储,REST + WebSocket 双协议

+ +
+
9
K线周期
+
5500+
A股标的
+
2
数据协议
+
全天
自动盘后拉取
+
+ +
+ K线周期 + 股票查询 + K线查询 + 实时行情 + 交易日历 + WebSocket + CLI 工具 +
+ + +

支持的 K 线周期

+
+
1m
1 分钟
+
5m
5 分钟
+
15m
15 分钟
+
30m
30 分钟
+
1h
1 小时
+
2h
2 小时(推导)
+
1d
日线
+
1w
周线
+
1M
月线
+
+ +
+ 注意:日线/周线/月线可回填全部历史数据;分钟线 API 仅保留近 1-3 个月,需通过每日盘后拉取积累本地数据; + 2 小时线从 1 小时线实时推导,无需单独存储。 +
+ + +

股票查询 GET

+ +
+

GET /api/v1/stocks

+

分页查询股票列表,可按交易所、板块筛选

+
+ 参数: exchange=SH|SZ|BJ   market=主板|创业板|科创板   limit=100   offset=0 +
+
+ +
+

GET /api/v1/stocks/{ts_code}

+

单只股票基本信息 — 例: /api/v1/stocks/000001.SZ

+
+ +
+

GET /api/v1/stocks/search?q=平安

+

按名称或代码模糊搜索,最多返回 50 条

+
+ +
# 搜索平安
+curl "http://localhost:8000/api/v1/stocks/search?q=平安"
+
+# 查看深交所主板股票
+curl "http://localhost:8000/api/v1/stocks?exchange=SZ&market=主板&limit=20"
+ + +

K 线查询 GET POST

+ +
+

GET /api/v1/klines/{freq}?ts_code=000001.SZ&start_date=2026-01-01&end_date=2026-05-16

+

查询单只股票指定周期的 K 线数据

+
+ 参数: freq=1m|5m|15m|30m|1h|2h|1d|1w|1M   ts_code   start_date   end_date   limit=10000 +
+
+ +
+

POST /api/v1/klines/{freq}/batch

+

批量查询多只股票 K 线

+
Body: {"codes": ["000001.SZ","600000.SH"], "start_date": "2026-01-01", "end_date": "2026-05-16"}
+
+ +
+

GET /api/v1/klines/{freq}/latest?ts_code=000001.SZ

+

最新交易日数据(ts_code 可选,不传返回全量)

+
+ +
+

GET /api/v1/klines/available-freqs

+

返回所有支持的频率及其分类

+
+ +
# 平安银行 2026 年日线
+curl "http://localhost:8000/api/v1/klines/1d?ts_code=000001.SZ&start_date=2026-01-01"
+
+# 批量查 2h K 线
+curl -X POST "http://localhost:8000/api/v1/klines/2h/batch" \
+  -H "Content-Type: application/json" \
+  -d '{"codes": ["000001.SZ","600000.SH"], "start_date": "2026-05-01"}'
+ + +

实时行情 GET

+ +
+

GET /api/v1/realtime/spot?codes=000001.SZ,600000.SH

+

获取指定股票实时快照(最新价、涨跌幅、成交量等)。不传 codes 返回全市场数据

+
+ +
+

GET /api/v1/realtime/market-state

+

当前市场状态:trading / lunch_break / closed / pre_open / closing_auction

+
+ +
curl "http://localhost:8000/api/v1/realtime/spot?codes=000001.SZ,600519.SH"
+curl "http://localhost:8000/api/v1/realtime/market-state"
+ + +

交易日历 GET

+ +
+

GET /api/v1/calendar/trading-days?start=2026-01-01&end=2026-05-16

+

查询日期范围内的交易日

+
+ +
+

GET /api/v1/calendar/is-trading-day?date=2026-01-15

+

判断指定日期是否为交易日

+
+ + +

实时推送 WebSocket

+ +
+

WS /ws/realtime

+

交易时段(9:25-15:05)每 5 秒推送订阅股票的实时行情

+
+ +
+ WebSocket 连接后需先发送订阅命令才会收到数据推送。支持按股票代码粒度订阅/退订。 +
+ +
# 使用 websocat 测试:
+websocat ws://localhost:8000/ws/realtime
+
+# 连接后发送订阅:
+{"action": "subscribe", "codes": ["000001.SZ", "600519.SH"]}
+
+# 退订:
+{"action": "unsubscribe", "codes": ["000001.SZ"]}
+
+# 取消全部订阅:
+{"action": "unsubscribe_all"}
+ +
# 服务器推送的消息格式:
+
+# 行情数据
+{"type": "spot", "data": {
+  "000001.SZ": {
+    "price": 12.34, "change": 0.12, "pct_chg": 0.98,
+    "volume": 12345678, "amount": 152345678.9,
+    "high": 12.45, "low": 12.10, "open": 12.20,
+    "pre_close": 12.22, "name": "平安银行"
+  }
+}}
+
+# 市场状态
+{"type": "market_state", "state": "trading", "timestamp": "2026-05-16T10:30:00+08:00"}
+
+# 心跳 (每 30 秒)
+{"type": "heartbeat", "timestamp": "2026-05-16T10:30:00+08:00", "connections": 3}
+
+# 确认订阅
+{"type": "subscribed", "codes": ["000001.SZ", "600519.SH"]}
+ + +

CLI 命令行工具

+ +
# 初始化数据库
+ashare-dp backfill init
+
+# 回填全部历史日线/周线/月线
+ashare-dp backfill daily
+
+# 回填近 30 天分钟数据
+ashare-dp backfill minute --days 30
+
+# 启动 API 服务
+ashare-dp serve start --port 8000
+
+# 命令行查询
+ashare-dp query kline 1d 000001.SZ --start 2026-01-01
+ashare-dp query latest --freq 1d --ts-code 000001.SZ
+ashare-dp query stats
+ +
+ + +
+ + +""" + + +@asynccontextmanager +async def lifespan(app: FastAPI): + """Startup and shutdown lifecycle.""" + logger.info("Starting A-Share Data Platform...") + db = get_db() + db.connect() + for ddl in DDL_STATEMENTS: + try: + db.execute(ddl) + except Exception as e: + logger.warning(f"DDL warning: {e}") + logger.info("Database initialized") + + # Start realtime poller + try: + from ashare_dp.data.realtime import poller + await poller.start() + app.state.poller = poller + logger.info("Realtime poller started") + except Exception as e: + logger.warning(f"Realtime poller not started: {e}") + app.state.poller = None + + # Start scheduler + try: + from ashare_dp.scheduler.scheduler import Scheduler + from ashare_dp.data.akshare_client import AKShareClient + scheduler = Scheduler(client=AKShareClient()) + scheduler.start() + app.state.scheduler = scheduler + logger.info("Scheduler started") + except Exception as e: + logger.warning(f"Scheduler not started: {e}") + app.state.scheduler = None + + yield + + # Shutdown + if app.state.poller: + try: + await app.state.poller.stop() + except Exception: + pass + if app.state.scheduler: + try: + app.state.scheduler.shutdown() + except Exception: + pass + db.close() + logger.info("A-Share Data Platform stopped") + + +def create_app() -> FastAPI: + """Create and configure the FastAPI application.""" + app = FastAPI( + title="A-Share Data Platform", + description="REST and WebSocket APIs for A-share market data", + version="0.1.0", + lifespan=lifespan, + ) + + # Root: documentation page + @app.get("/", response_class=HTMLResponse) + async def root(): + return DOCS_HTML + + @app.get("/health") + async def health(): + """Basic health check.""" + db = get_db() + repo = KLineRepository() + return { + "status": "ok", + "db_path": db.db_path, + "stocks": db.query("SELECT count(*) FROM stock_info")[0][0], + "trading_days": db.query("SELECT count(*) FROM trading_calendar")[0][0], + "daily_records": repo.count_records(Freq.d1), + } + + @app.get("/stats") + async def stats(): + """Full DB statistics.""" + db = get_db() + repo = KLineRepository() + from ashare_dp.core.models import Freq + freq_stats = {} + for f in [Freq.d1, Freq.w1, Freq.M1, Freq.m1, Freq.m5, Freq.m15, Freq.m30, Freq.h1]: + dr = repo.get_date_range(f) + freq_stats[f.value] = { + "records": repo.count_records(f), + "start_date": dr[0].isoformat() if dr[0] else None, + "end_date": dr[1].isoformat() if dr[1] else None, + } + return { + "stocks": db.query("SELECT count(*) FROM stock_info")[0][0], + "trading_days": db.query("SELECT count(*) FROM trading_calendar")[0][0], + "frequencies": freq_stats, + } + + # Register routers + app.include_router(stocks.router, prefix="/api/v1") + app.include_router(kline.router, prefix="/api/v1") + app.include_router(realtime.router, prefix="/api/v1") + app.include_router(calendar.router, prefix="/api/v1") + app.include_router(ws_router) + + return app diff --git a/src/ashare_dp/api/deps.py b/src/ashare_dp/api/deps.py new file mode 100644 index 0000000..57f0479 --- /dev/null +++ b/src/ashare_dp/api/deps.py @@ -0,0 +1,9 @@ +"""FastAPI dependency injection.""" + +from ashare_dp.storage.database import get_db +from ashare_dp.storage.repository import KLineRepository + + +def get_repo() -> KLineRepository: + """Get the K-line repository instance.""" + return KLineRepository() diff --git a/src/ashare_dp/api/routers/__init__.py b/src/ashare_dp/api/routers/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/ashare_dp/api/routers/calendar.py b/src/ashare_dp/api/routers/calendar.py new file mode 100644 index 0000000..b7f788d --- /dev/null +++ b/src/ashare_dp/api/routers/calendar.py @@ -0,0 +1,64 @@ +"""Trading calendar REST endpoints.""" + +from datetime import date, datetime + +from fastapi import APIRouter, HTTPException, Query +from loguru import logger + +from ashare_dp.core.calendar import calendar, determine_market_state, is_trading_time + +router = APIRouter(prefix="/calendar", tags=["calendar"]) + + +@router.get("/trading-days") +async def trading_days( + start: str = Query(..., description="Start date YYYY-MM-DD"), + end: str = Query(..., description="End date YYYY-MM-DD"), +): + """List trading days in a date range (inclusive).""" + try: + sd = date.fromisoformat(start) + ed = date.fromisoformat(end) + except ValueError as e: + raise HTTPException(status_code=400, detail=f"Invalid date format: {e}") + + days = calendar.get_trading_days(sd, ed) + return { + "start": start, + "end": end, + "count": len(days), + "trading_days": [d.isoformat() for d in days], + } + + +@router.get("/is-trading-day") +async def check_trading_day( + d: str = Query(..., alias="date", description="Date YYYY-MM-DD"), +): + """Check if a specific date is a trading day.""" + try: + dt = date.fromisoformat(d) + except ValueError as e: + raise HTTPException(status_code=400, detail=f"Invalid date: {e}") + + return { + "date": d, + "is_trading_day": calendar.is_trading_day(dt), + } + + +@router.get("/next-trading-day") +async def next_trading_day( + d: str = Query(..., alias="date", description="Date YYYY-MM-DD"), +): + """Get the next trading day on or after the given date.""" + try: + dt = date.fromisoformat(d) + except ValueError as e: + raise HTTPException(status_code=400, detail=f"Invalid date: {e}") + + nd = calendar.next_trading_day(dt) + return { + "date": d, + "next_trading_day": nd.isoformat(), + } diff --git a/src/ashare_dp/api/routers/kline.py b/src/ashare_dp/api/routers/kline.py new file mode 100644 index 0000000..53f59c4 --- /dev/null +++ b/src/ashare_dp/api/routers/kline.py @@ -0,0 +1,140 @@ +"""K-line query REST endpoints.""" + +from datetime import date, datetime +from typing import Optional + +from fastapi import APIRouter, Depends, HTTPException, Query +from loguru import logger + +from ashare_dp.api.deps import get_repo +from ashare_dp.core.models import BACKFILLABLE_FREQS, DERIVED_FREQS, INTRADAY_FREQS, Freq +from ashare_dp.storage.repository import KLineRepository + +router = APIRouter(prefix="/klines", tags=["klines"]) + +ALL_FREQS = BACKFILLABLE_FREQS + INTRADAY_FREQS + DERIVED_FREQS +FREQ_VALUES = [f.value for f in ALL_FREQS] + + +def _df_to_rows(df): + """Convert DataFrame to list of dicts, handling date/time serialization.""" + if df is None or df.empty: + return [] + df = df.copy() + for col in df.columns: + if df[col].dtype.name.startswith("datetime"): + df[col] = df[col].apply(lambda x: x.isoformat() if hasattr(x, "isoformat") else x) + elif df[col].dtype.name == "object": + # Handle date objects + pass + return df.to_dict(orient="records") + + +@router.get("/available-freqs") +async def available_freqs(): + """List all supported K-line frequencies.""" + return { + "frequencies": FREQ_VALUES, + "backfillable": [f.value for f in BACKFILLABLE_FREQS], + "intraday": [f.value for f in INTRADAY_FREQS], + "derived": [f.value for f in DERIVED_FREQS], + } + + +@router.get("/{freq}") +async def query_klines( + freq: str, + repo: KLineRepository = Depends(get_repo), + ts_code: Optional[str] = Query(None, description="Stock code, e.g. '000001.SZ'"), + start_date: Optional[str] = Query(None, description="Start date (YYYY-MM-DD)"), + end_date: Optional[str] = Query(None, description="End date (YYYY-MM-DD)"), + limit: int = Query(10000, ge=1, le=100000), + offset: int = Query(0, ge=0), +): + """Query K-line data by frequency. + + Supported frequencies: 1m, 5m, 15m, 30m, 1h, 2h, 1d, 1w, 1M + """ + if freq not in FREQ_VALUES: + raise HTTPException( + status_code=400, + detail=f"Invalid frequency: {freq}. Valid: {FREQ_VALUES}", + ) + + freq_enum = Freq(freq) + + # Parse dates + sd = date.fromisoformat(start_date) if start_date else None + ed = date.fromisoformat(end_date) if end_date else None + + df = repo.read_klines( + freq=freq_enum, + ts_code=ts_code, + start_date=sd, + end_date=ed, + limit=limit, + offset=offset, + ) + + rows = _df_to_rows(df) + return {"freq": freq, "count": len(rows), "items": rows} + + +@router.post("/{freq}/batch") +async def batch_query_klines( + freq: str, + body: dict, + repo: KLineRepository = Depends(get_repo), +): + """Query K-line data for multiple stocks in a single request. + + Body: {"codes": ["000001.SZ", "600000.SH"], "start_date": "...", "end_date": "..."} + """ + if freq not in FREQ_VALUES: + raise HTTPException( + status_code=400, + detail=f"Invalid frequency: {freq}. Valid: {FREQ_VALUES}", + ) + + freq_enum = Freq(freq) + codes = body.get("codes", []) + if not codes: + raise HTTPException(status_code=400, detail="codes list is required") + + start_date = body.get("start_date") + end_date = body.get("end_date") + sd = date.fromisoformat(start_date) if start_date else None + ed = date.fromisoformat(end_date) if end_date else None + + results = {} + for code in codes: + df = repo.read_klines( + freq=freq_enum, + ts_code=code, + start_date=sd, + end_date=ed, + limit=body.get("limit", 10000), + offset=body.get("offset", 0), + ) + results[code] = _df_to_rows(df) + + return {"freq": freq, "results": results} + + +@router.get("/{freq}/latest") +async def latest_klines( + freq: str, + repo: KLineRepository = Depends(get_repo), + ts_code: Optional[str] = Query(None, description="Stock code (optional)"), +): + """Get the latest K-line data for the most recent trading day.""" + if freq not in FREQ_VALUES: + raise HTTPException( + status_code=400, + detail=f"Invalid frequency: {freq}. Valid: {FREQ_VALUES}", + ) + + freq_enum = Freq(freq) + df = repo.get_latest(freq=freq_enum, ts_code=ts_code) + rows = _df_to_rows(df) + return {"freq": freq, "count": len(rows), "items": rows} diff --git a/src/ashare_dp/api/routers/realtime.py b/src/ashare_dp/api/routers/realtime.py new file mode 100644 index 0000000..40c5a05 --- /dev/null +++ b/src/ashare_dp/api/routers/realtime.py @@ -0,0 +1,88 @@ +"""Real-time spot and market state REST endpoints.""" + +from datetime import datetime +from typing import Optional + +from fastapi import APIRouter, HTTPException, Query +from loguru import logger + +from ashare_dp.core.calendar import BEIJING_TZ, determine_market_state +from ashare_dp.data.akshare_client import AKShareClient + +router = APIRouter(prefix="/realtime", tags=["realtime"]) + +_client = AKShareClient() + + +@router.get("/spot") +async def get_spot( + codes: Optional[str] = Query(None, description="Comma-separated stock codes (e.g. '000001.SZ,600000.SH')"), +): + """Get real-time spot data. If codes not specified, returns all stocks.""" + try: + df = _client.get_spot() + except Exception as e: + raise HTTPException(status_code=502, detail=f"Failed to fetch spot data: {e}") + + if df.empty: + return {"timestamp": datetime.now(BEIJING_TZ).isoformat(), "count": 0, "items": []} + + # Map Chinese column names + col_map = { + "代码": "code", + "名称": "name", + "最新价": "price", + "涨跌额": "change", + "涨跌幅": "pct_chg", + "成交量": "volume", + "成交额": "amount", + "最高": "high", + "最低": "low", + "今开": "open", + "昨收": "pre_close", + } + df = df.rename(columns={k: v for k, v in col_map.items() if k in df.columns}) + + # Build ts_code from code + if "code" in df.columns: + def _to_ts_code(c): + c = str(c).zfill(6) + if c.startswith(("6", "9")): + return f"{c}.SH" + elif c.startswith(("8", "4")): + return f"{c}.BJ" + return f"{c}.SZ" + df["ts_code"] = df["code"].apply(_to_ts_code) + + # Filter by requested codes + if codes: + code_set = set(c.strip() for c in codes.split(",")) + if "ts_code" in df.columns: + df = df[df["ts_code"].isin(code_set)] + + # Select relevant columns + out_cols = [c for c in ["ts_code", "code", "name", "price", "change", "pct_chg", + "volume", "amount", "high", "low", "open", "pre_close"] + if c in df.columns] + df = df[out_cols] + + # Convert to dicts + items = df.to_dict(orient="records") + return { + "timestamp": datetime.now(BEIJING_TZ).isoformat(), + "count": len(items), + "items": items, + } + + +@router.get("/market-state") +async def market_state(): + """Get current A-share market state.""" + now = datetime.now(BEIJING_TZ) + state = determine_market_state(now) + return { + "timestamp": now.isoformat(), + "state": state, + "is_trading": state == "trading", + "weekday": now.weekday(), + } diff --git a/src/ashare_dp/api/routers/stocks.py b/src/ashare_dp/api/routers/stocks.py new file mode 100644 index 0000000..4b26166 --- /dev/null +++ b/src/ashare_dp/api/routers/stocks.py @@ -0,0 +1,96 @@ +"""Stock info REST endpoints.""" + +from datetime import date +from typing import Optional + +from fastapi import APIRouter, HTTPException, Query +from loguru import logger + +from ashare_dp.storage.database import get_db + +router = APIRouter(prefix="/stocks", tags=["stocks"]) + + +@router.get("") +async def list_stocks( + exchange: Optional[str] = Query(None, description="Exchange: SH, SZ, BJ"), + market: Optional[str] = Query(None, description="Market: 主板, 创业板, 科创板, 北交所"), + limit: int = Query(100, ge=1, le=1000), + offset: int = Query(0, ge=0), +): + """List stocks with optional filters.""" + db = get_db() + conditions = ["1=1"] + params = [] + + if exchange: + conditions.append("exchange = ?") + params.append(exchange.upper()) + if market: + conditions.append("market = ?") + params.append(market) + + where = " AND ".join(conditions) + rows = db.query( + f"SELECT * FROM stock_info WHERE {where} ORDER BY ts_code LIMIT ? OFFSET ?", + tuple(params) + (limit, offset), + ) + total = db.query( + f"SELECT count(*) FROM stock_info WHERE {where}", + tuple(params), + )[0][0] + + cols = ["ts_code", "symbol", "name", "exchange", "area", "industry", "list_date", "delist_date", "market", "updated_at"] + items = [] + for row in rows: + item = {} + for i, col in enumerate(cols): + val = row[i] if i < len(row) else None + if isinstance(val, date): + val = val.isoformat() + item[col] = val + items.append(item) + + return {"total": total, "offset": offset, "limit": limit, "items": items} + + +@router.get("/search") +async def search_stocks(q: str = Query(..., min_length=1, description="Search query")): + """Fuzzy search stocks by name or code.""" + db = get_db() + rows = db.query( + "SELECT * FROM stock_info WHERE name LIKE ? OR symbol LIKE ? OR ts_code LIKE ? LIMIT 50", + (f"%{q}%", f"%{q}%", f"%{q}%"), + ) + cols = ["ts_code", "symbol", "name", "exchange", "area", "industry", "list_date", "delist_date", "market", "updated_at"] + items = [] + for row in rows: + item = {} + for i, col in enumerate(cols): + val = row[i] if i < len(row) else None + if isinstance(val, date): + val = val.isoformat() + item[col] = val + items.append(item) + return {"query": q, "count": len(items), "items": items} + + +@router.get("/{ts_code}") +async def get_stock(ts_code: str): + """Get single stock info by ts_code (e.g. '000001.SZ').""" + db = get_db() + row = db.query( + "SELECT * FROM stock_info WHERE ts_code = ?", + (ts_code,), + ) + if not row: + raise HTTPException(status_code=404, detail=f"Stock not found: {ts_code}") + row = row[0] + cols = ["ts_code", "symbol", "name", "exchange", "area", "industry", "list_date", "delist_date", "market", "updated_at"] + item = {} + for i, col in enumerate(cols): + val = row[i] if i < len(row) else None + if isinstance(val, date): + val = val.isoformat() + item[col] = val + return item diff --git a/src/ashare_dp/api/websocket/__init__.py b/src/ashare_dp/api/websocket/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/ashare_dp/api/websocket/handlers.py b/src/ashare_dp/api/websocket/handlers.py new file mode 100644 index 0000000..fa8cf29 --- /dev/null +++ b/src/ashare_dp/api/websocket/handlers.py @@ -0,0 +1,78 @@ +"""WebSocket message handlers and route registration.""" + +from __future__ import annotations + +from datetime import datetime + +from fastapi import APIRouter, WebSocket, WebSocketDisconnect +from loguru import logger + +from ashare_dp.api.websocket.manager import manager +from ashare_dp.core.calendar import BEIJING_TZ, determine_market_state + +router = APIRouter() + + +@router.websocket("/ws/realtime") +async def websocket_realtime(websocket: WebSocket): + """WebSocket endpoint for real-time market data. + + Client messages: + {"action": "subscribe", "codes": ["000001.SZ", "600000.SH"]} + {"action": "unsubscribe", "codes": ["000001.SZ"]} + {"action": "unsubscribe_all"} + + Server messages: + {"type": "spot", "data": {"000001.SZ": {...}}} + {"type": "market_state", "state": "...", "timestamp": "..."} + {"type": "heartbeat", "timestamp": "..."} + {"type": "error", "message": "..."} + """ + client_id = await manager.connect(websocket) + + try: + while True: + data = await websocket.receive_json() + action = data.get("action") + + if action == "subscribe": + codes = data.get("codes", []) + if not codes: + await manager.send_to_client(client_id, { + "type": "error", + "message": "codes list is required for subscribe", + }) + continue + await manager.subscribe(client_id, codes) + await manager.send_to_client(client_id, { + "type": "subscribed", + "codes": list(await manager.get_client_codes(client_id)), + }) + + elif action == "unsubscribe": + codes = data.get("codes", []) + await manager.unsubscribe(client_id, codes) + await manager.send_to_client(client_id, { + "type": "unsubscribed", + "codes": codes, + "remaining": list(await manager.get_client_codes(client_id)), + }) + + elif action == "unsubscribe_all": + await manager.unsubscribe_all(client_id) + await manager.send_to_client(client_id, { + "type": "unsubscribed_all", + }) + + else: + await manager.send_to_client(client_id, { + "type": "error", + "message": f"Unknown action: {action}. Valid: subscribe, unsubscribe, unsubscribe_all", + }) + + except WebSocketDisconnect: + pass + except Exception as e: + logger.error(f"WS error for client {client_id}: {e}") + finally: + await manager.disconnect(client_id) diff --git a/src/ashare_dp/api/websocket/manager.py b/src/ashare_dp/api/websocket/manager.py new file mode 100644 index 0000000..472867a --- /dev/null +++ b/src/ashare_dp/api/websocket/manager.py @@ -0,0 +1,152 @@ +"""WebSocket connection and subscription manager.""" + +from __future__ import annotations + +import asyncio +import uuid + +from fastapi import WebSocket +from loguru import logger + + +class ConnectionManager: + """Manages WebSocket connections and per-client subscriptions. + + Tracks: + - Active connections (client_id -> WebSocket) + - Per-client subscriptions (client_id -> set of ts_codes) + - Reverse index (ts_code -> set of client_ids) for fast lookup + """ + + def __init__(self): + self._connections: dict[str, WebSocket] = {} + self._subscriptions: dict[str, set[str]] = {} + self._code_subscribers: dict[str, set[str]] = {} + self._lock = asyncio.Lock() + + async def connect(self, websocket: WebSocket) -> str: + """Accept a new WebSocket connection and return a client_id.""" + await websocket.accept() + client_id = str(uuid.uuid4())[:8] + async with self._lock: + self._connections[client_id] = websocket + self._subscriptions[client_id] = set() + logger.info(f"WS client connected: {client_id}") + return client_id + + async def disconnect(self, client_id: str): + """Remove a client and all its subscriptions.""" + async with self._lock: + if client_id in self._subscriptions: + # Remove from reverse index + for code in self._subscriptions[client_id]: + if code in self._code_subscribers: + self._code_subscribers[code].discard(client_id) + if not self._code_subscribers[code]: + del self._code_subscribers[code] + del self._subscriptions[client_id] + self._connections.pop(client_id, None) + logger.info(f"WS client disconnected: {client_id}") + + async def subscribe(self, client_id: str, codes: list[str]): + """Subscribe a client to specific stock codes.""" + async with self._lock: + if client_id not in self._subscriptions: + return + for code in codes: + self._subscriptions[client_id].add(code) + if code not in self._code_subscribers: + self._code_subscribers[code] = set() + self._code_subscribers[code].add(client_id) + logger.debug(f"Client {client_id} subscribed to {len(codes)} codes") + + async def unsubscribe(self, client_id: str, codes: list[str]): + """Unsubscribe a client from specific codes.""" + async with self._lock: + if client_id not in self._subscriptions: + return + for code in codes: + self._subscriptions[client_id].discard(code) + if code in self._code_subscribers: + self._code_subscribers[code].discard(client_id) + if not self._code_subscribers[code]: + del self._code_subscribers[code] + + async def unsubscribe_all(self, client_id: str): + """Unsubscribe a client from all codes.""" + async with self._lock: + if client_id not in self._subscriptions: + return + codes = list(self._subscriptions[client_id]) + await self.unsubscribe(client_id, codes) + + async def get_all_subscribed_codes(self) -> set[str]: + """Get the union of all codes any client is subscribed to.""" + async with self._lock: + return set(self._code_subscribers.keys()) + + async def get_client_codes(self, client_id: str) -> set[str]: + """Get codes a specific client is subscribed to.""" + async with self._lock: + return self._subscriptions.get(client_id, set()).copy() + + async def send_to_client(self, client_id: str, message: dict): + """Send a JSON message to a specific client.""" + ws = self._connections.get(client_id) + if ws is None: + return + try: + await ws.send_json(message) + except Exception: + await self.disconnect(client_id) + + async def broadcast(self, message: dict, client_ids: set[str] | None = None): + """Broadcast a message to specific clients, or all if None.""" + if client_ids is None: + client_ids = set(self._connections.keys()) + + tasks = [] + for cid in client_ids: + tasks.append(self.send_to_client(cid, message)) + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) + + async def broadcast_filtered(self, spot_data: dict[str, dict]): + """Broadcast spot data, filtering per-client based on subscriptions. + + Args: + spot_data: Dict mapping ts_code -> spot data dict. + """ + # Snapshot connections and subscriptions under lock + async with self._lock: + client_codes = { + cid: codes.copy() + for cid, codes in self._subscriptions.items() + if cid in self._connections + } + + tasks = [] + for cid, codes in client_codes.items(): + if not codes: + continue + # Filter to only subscribed codes for this client + filtered = {c: spot_data[c] for c in codes if c in spot_data} + if filtered: + tasks.append(self.send_to_client(cid, { + "type": "spot", + "data": filtered, + })) + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) + + @property + def active_connections(self) -> int: + return len(self._connections) + + @property + def active_subscriptions(self) -> int: + return sum(len(v) for v in self._subscriptions.values()) + + +# Global manager instance +manager = ConnectionManager() diff --git a/src/ashare_dp/cli/__init__.py b/src/ashare_dp/cli/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/ashare_dp/cli/backfill_cmd.py b/src/ashare_dp/cli/backfill_cmd.py new file mode 100644 index 0000000..713a594 --- /dev/null +++ b/src/ashare_dp/cli/backfill_cmd.py @@ -0,0 +1,92 @@ +"""Backfill CLI subcommands.""" + +from __future__ import annotations + +import asyncio +from datetime import date, datetime + +import typer +from loguru import logger + +from ashare_dp.core.models import BACKFILLABLE_FREQS, INTRADAY_FREQS +from ashare_dp.data.akshare_client import AKShareClient +from ashare_dp.data.backfill import BackfillPipeline +from ashare_dp.storage.database import get_db + +backfill_app = typer.Typer() + + +@backfill_app.command("init") +def init_db(): + """Initialize database schema.""" + db = get_db() + db.connect() + pipeline = BackfillPipeline() + pipeline.init_db() + pipeline.load_stock_list() + pipeline.load_trading_calendar() + typer.echo("Database initialized with stock list and trading calendar") + + +@backfill_app.command("daily") +def backfill_daily( + start: str = typer.Option("19900101", help="Start date YYYYMMDD"), + end: str = typer.Option(None, help="End date YYYYMMDD (default: today)"), + workers: int = typer.Option(10, help="Number of worker threads"), + symbols: str = typer.Option(None, help="Comma-separated stock symbols (default: all)"), +): + """Backfill daily/weekly/monthly K-line data.""" + db = get_db() + db.connect() + + start_date = datetime.strptime(start, "%Y%m%d").date() + end_date = datetime.strptime(end, "%Y%m%d").date() if end else date.today() + + sym_list = [s.strip() for s in symbols.split(",")] if symbols else None + + pipeline = BackfillPipeline(max_workers=workers) + pipeline.init_db() + pipeline.load_stock_list() + pipeline.load_trading_calendar() + + results = pipeline.backfill_daily_weekly_monthly( + symbols=sym_list, + start_date=start_date, + end_date=end_date, + ) + + typer.echo(f"\nBackfill complete:") + for freq, stats in results.items(): + typer.echo( + f" {freq}: {stats['records']} records, " + f"{stats['completed']} stocks ok, {stats['failed']} failed" + ) + + +@backfill_app.command("minute") +def backfill_minute( + days: int = typer.Option(30, help="Number of calendar days to look back"), + workers: int = typer.Option(5, help="Number of worker threads"), + symbols: str = typer.Option(None, help="Comma-separated stock symbols (default: all)"), +): + """Backfill recent minute K-line data (limited API history).""" + db = get_db() + db.connect() + + sym_list = [s.strip() for s in symbols.split(",")] if symbols else None + + pipeline = BackfillPipeline(max_workers=workers) + pipeline.init_db() + pipeline.load_stock_list() + + results = pipeline.backfill_minute( + symbols=sym_list, + days_back=days, + ) + + typer.echo(f"\nMinute backfill complete ({days} day lookback):") + for freq, stats in results.items(): + typer.echo( + f" {freq}: {stats['records']} records, " + f"{stats['completed']} stocks ok, {stats['failed']} failed" + ) diff --git a/src/ashare_dp/cli/main.py b/src/ashare_dp/cli/main.py new file mode 100644 index 0000000..de4cb5b --- /dev/null +++ b/src/ashare_dp/cli/main.py @@ -0,0 +1,29 @@ +"""Typer CLI entry point for ashare-dp.""" + +from __future__ import annotations + +import typer + +from ashare_dp.cli.backfill_cmd import backfill_app +from ashare_dp.cli.serve_cmd import serve_app +from ashare_dp.cli.query_cmd import query_app + +app = typer.Typer( + name="ashare-dp", + help="A-Share Data Platform CLI", +) + +app.add_typer(backfill_app, name="backfill", help="Historical data backfill") +app.add_typer(serve_app, name="serve", help="Start API server") +app.add_typer(query_app, name="query", help="Ad-hoc data queries") + + +@app.command() +def version(): + """Show version.""" + from ashare_dp import __version__ + typer.echo(f"ashare-dp v{__version__}") + + +if __name__ == "__main__": + app() diff --git a/src/ashare_dp/cli/query_cmd.py b/src/ashare_dp/cli/query_cmd.py new file mode 100644 index 0000000..7dcac1e --- /dev/null +++ b/src/ashare_dp/cli/query_cmd.py @@ -0,0 +1,113 @@ +"""Query CLI subcommands for ad-hoc data queries.""" + +from __future__ import annotations + +from datetime import date, datetime + +import typer + +from ashare_dp.core.models import Freq +from ashare_dp.storage.database import get_db +from ashare_dp.storage.repository import KLineRepository + +query_app = typer.Typer() + + +@query_app.command("kline") +def query_kline( + freq: str = typer.Argument(..., help="Frequency: 1m,5m,15m,30m,1h,2h,1d,1w,1M"), + ts_code: str = typer.Argument(..., help="Stock code, e.g. 000001.SZ"), + start: str = typer.Option(None, help="Start date YYYY-MM-DD"), + end: str = typer.Option(None, help="End date YYYY-MM-DD"), + limit: int = typer.Option(100, help="Max records"), +): + """Query K-line data from the command line.""" + db = get_db() + db.connect() + + freq_enum = Freq(freq) + repo = KLineRepository() + + sd = date.fromisoformat(start) if start else None + ed = date.fromisoformat(end) if end else None + + df = repo.read_klines(freq=freq_enum, ts_code=ts_code, start_date=sd, end_date=ed, limit=limit) + + if df.empty: + typer.echo("No data found") + return + + typer.echo(f"\n{freq} K-line for {ts_code}:") + typer.echo(df.to_string(index=False)) + typer.echo(f"\n{len(df)} records") + + +@query_app.command("latest") +def query_latest( + freq: str = typer.Option("1d", help="Frequency"), + ts_code: str = typer.Option(None, help="Stock code (optional)"), +): + """Show latest K-line data.""" + db = get_db() + db.connect() + + freq_enum = Freq(freq) + repo = KLineRepository() + df = repo.get_latest(freq=freq_enum, ts_code=ts_code) + + if df.empty: + typer.echo("No data found") + return + + typer.echo(f"\nLatest {freq} K-line:") + typer.echo(df.to_string(index=False)) + typer.echo(f"\n{len(df)} records") + + +@query_app.command("stocks") +def query_stocks( + exchange: str = typer.Option(None, help="Exchange: SH, SZ, BJ"), + limit: int = typer.Option(50, help="Max records"), +): + """List stocks.""" + db = get_db() + db.connect() + + if exchange: + rows = db.query( + "SELECT ts_code, symbol, name, exchange, market, list_date FROM stock_info WHERE exchange = ? LIMIT ?", + (exchange.upper(), limit), + ) + else: + rows = db.query( + "SELECT ts_code, symbol, name, exchange, market, list_date FROM stock_info LIMIT ?", + (limit,), + ) + + typer.echo(f"\n{'ts_code':<12} {'symbol':<8} {'name':<12} {'exchange':<8} {'market':<10} {'list_date'}") + typer.echo("-" * 60) + for row in rows: + ts, sym, name, ex, mkt, ld = row + ld_str = str(ld) if ld else "" + typer.echo(f"{ts:<12} {sym:<8} {name:<12} {ex:<8} {mkt or '':<10} {ld_str}") + + +@query_app.command("stats") +def query_stats(): + """Show database statistics.""" + db = get_db() + db.connect() + repo = KLineRepository() + + typer.echo("\nDatabase Statistics:") + typer.echo("-" * 40) + + # Stock count + n = db.query("SELECT count(*) FROM stock_info")[0][0] + typer.echo(f" Stocks: {n}") + + # Record count and date range per frequency + for freq_repr in [Freq.d1, Freq.w1, Freq.M1, Freq.h1, Freq.m5, Freq.m1]: + count = repo.count_records(freq_repr) + dr = repo.get_date_range(freq_repr) + typer.echo(f" {freq_repr.value}: {count} records, range {dr[0]} ~ {dr[1]}") diff --git a/src/ashare_dp/cli/serve_cmd.py b/src/ashare_dp/cli/serve_cmd.py new file mode 100644 index 0000000..8bc1797 --- /dev/null +++ b/src/ashare_dp/cli/serve_cmd.py @@ -0,0 +1,39 @@ +"""Serve CLI subcommand: start the API server.""" + +from __future__ import annotations + +import asyncio + +import typer +import uvicorn +from loguru import logger + +from ashare_dp.config import Settings + +settings = Settings() +serve_app = typer.Typer() + + +@serve_app.command("start") +def serve( + host: str = typer.Option(None, help="Bind address"), + port: int = typer.Option(None, help="Bind port"), + reload: bool = typer.Option(False, help="Enable auto-reload (dev mode)"), +): + """Start the API server with scheduler.""" + h = host or settings.api_host + p = port or settings.api_port + + logger.info(f"Starting API server on {h}:{p}") + + # Start the realtime poller in a background thread via the app lifespan + # The FastAPI lifespan handles DB init, scheduler start, and poller start + + uvicorn.run( + "ashare_dp.api.app:create_app", + host=h, + port=p, + reload=reload, + factory=True, + log_level=settings.log_level.lower(), + ) diff --git a/src/ashare_dp/config.py b/src/ashare_dp/config.py new file mode 100644 index 0000000..2a35cdc --- /dev/null +++ b/src/ashare_dp/config.py @@ -0,0 +1,40 @@ +from pydantic_settings import BaseSettings, SettingsConfigDict + + +class Settings(BaseSettings): + """Application settings loaded from .env and environment variables.""" + + model_config = SettingsConfigDict( + env_file=".env", + env_file_encoding="utf-8", + ) + + # Data paths + data_dir: str = "data" + duckdb_path: str = "data/duckdb/ashare.db" + + # API server + api_host: str = "0.0.0.0" + api_port: int = 8000 + + # Backfill + backfill_workers: int = 10 + + # AKShare + akshare_max_retries: int = 3 + akshare_retry_delay: float = 1.0 + akshare_backend: str = "auto" # auto, em (East Money), sina + + # Realtime + realtime_poll_interval: int = 5 + + # Logging + log_level: str = "INFO" + + @property + def parquet_dir(self) -> str: + return f"{self.data_dir}/parquet" + + @property + def duckdb_dir(self) -> str: + return f"{self.data_dir}/duckdb" diff --git a/src/ashare_dp/core/__init__.py b/src/ashare_dp/core/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/ashare_dp/core/calendar.py b/src/ashare_dp/core/calendar.py new file mode 100644 index 0000000..030266e --- /dev/null +++ b/src/ashare_dp/core/calendar.py @@ -0,0 +1,141 @@ +"""Trading calendar service for A-share market. + +Wraps akshare's trading calendar functions and provides query methods. +""" + +from __future__ import annotations + +from datetime import date, datetime, timedelta +from typing import Optional + +import pandas as pd +import pytz +from loguru import logger + +from ashare_dp.core.exceptions import DataSourceError + +BEIJING_TZ = pytz.timezone("Asia/Shanghai") + +# A-share market session times (Beijing time) +MORNING_OPEN = (9, 30) +MORNING_CLOSE = (11, 30) +AFTERNOON_OPEN = (13, 0) +AFTERNOON_CLOSE = (15, 0 ) +PRE_OPEN_START = (9, 25) +CLOSING_AUCTION_START = (14, 57) + + +class MarketState: + TRADING = "trading" + LUNCH_BREAK = "lunch_break" + CLOSED = "closed" + PRE_OPEN = "pre_open" + CLOSING_AUCTION = "closing_auction" + + +def determine_market_state(dt: datetime | None = None) -> str: + """Determine the current A-share market state. + + Args: + dt: A timezone-aware datetime in Beijing time. Defaults to now. + + Returns: + One of MarketState values. + """ + if dt is None: + dt = datetime.now(BEIJING_TZ) + elif dt.tzinfo is None: + dt = BEIJING_TZ.localize(dt) + + # Weekends are always closed + if dt.weekday() >= 5: + return MarketState.CLOSED + + t = (dt.hour, dt.minute) + + if t < PRE_OPEN_START: + return MarketState.CLOSED + if t < MORNING_OPEN: + return MarketState.PRE_OPEN + if t < MORNING_CLOSE: + return MarketState.TRADING + if t < AFTERNOON_OPEN: + return MarketState.LUNCH_BREAK + if t < CLOSING_AUCTION_START: + return MarketState.TRADING + if t < AFTERNOON_CLOSE: + return MarketState.CLOSING_AUCTION + return MarketState.CLOSED + + +def is_trading_time(dt: datetime | None = None) -> bool: + """Check if the market is currently in a trading session.""" + return determine_market_state(dt) == MarketState.TRADING + + +class TradingCalendar: + """Manages A-share trading calendar data.""" + + def __init__(self): + self._calendar: pd.DataFrame | None = None + + def load(self) -> pd.DataFrame: + """Load trading calendar from akshare. + + Returns a DataFrame with columns: trade_date, is_trading_day. + """ + try: + import akshare as ak + df = ak.tool_trade_date_hist_sina() + if "trade_date" in df.columns: + df["trade_date"] = pd.to_datetime(df["trade_date"]).dt.date + logger.info(f"Loaded trading calendar: {len(df)} days") + self._calendar = df + return df + except Exception as e: + raise DataSourceError(f"Failed to load trading calendar: {e}") from e + + def is_trading_day(self, d: date) -> bool: + """Check if a given date is a trading day.""" + if self._calendar is None: + self.load() + mask = self._calendar["trade_date"] == d + if mask.any(): + row = self._calendar[mask].iloc[0] + return bool(row.get("is_trading_day", True)) + # Fallback: weekday check + return d.weekday() < 5 + + def get_trading_days(self, start: date, end: date) -> list[date]: + """Get all trading days in a date range (inclusive).""" + if self._calendar is None: + self.load() + mask = ( + (self._calendar["trade_date"] >= start) + & (self._calendar["trade_date"] <= end) + ) + if "is_trading_day" in self._calendar.columns: + mask &= self._calendar["is_trading_day"] == 1 + days = self._calendar[mask]["trade_date"].tolist() + return [d if isinstance(d, date) else pd.Timestamp(d).date() for d in days] + + def next_trading_day(self, d: date) -> date: + """Get the next trading day on or after the given date.""" + if self._calendar is None: + self.load() + mask = (self._calendar["trade_date"] >= d) + if "is_trading_day" in self._calendar.columns: + mask &= self._calendar["is_trading_day"] == 1 + if mask.any(): + result = self._calendar[mask].iloc[0]["trade_date"] + return result if isinstance(result, date) else pd.Timestamp(result).date() + return d + timedelta(days=1) + + def today_is_trading_day(self) -> bool: + """Check if today (Beijing time) is a trading day.""" + today = datetime.now(BEIJING_TZ).date() + return self.is_trading_day(today) + + +# Global calendar instance +calendar = TradingCalendar() diff --git a/src/ashare_dp/core/exceptions.py b/src/ashare_dp/core/exceptions.py new file mode 100644 index 0000000..6c432cb --- /dev/null +++ b/src/ashare_dp/core/exceptions.py @@ -0,0 +1,17 @@ +"""Custom exceptions for the A-Share Data Platform.""" + + +class AShareDPError(Exception): + """Base exception for all application errors.""" + + +class DataSourceError(AShareDPError): + """Error fetching data from upstream source (akshare).""" + + +class StorageError(AShareDPError): + """Error reading from or writing to storage.""" + + +class QueryError(AShareDPError): + """Invalid query parameters.""" diff --git a/src/ashare_dp/core/models.py b/src/ashare_dp/core/models.py new file mode 100644 index 0000000..3e6b448 --- /dev/null +++ b/src/ashare_dp/core/models.py @@ -0,0 +1,87 @@ +from __future__ import annotations + +from enum import Enum +from typing import Optional + + +class Freq(str, Enum): + """K-line frequency.""" + + m1 = "1m" + m5 = "5m" + m15 = "15m" + m30 = "30m" + h1 = "1h" + h2 = "2h" # derived from 1h + d1 = "1d" + w1 = "1w" + M1 = "1M" + + @property + def is_intraday(self) -> bool: + """True for sub-daily frequencies.""" + return self in (Freq.m1, Freq.m5, Freq.m15, Freq.m30, Freq.h1, Freq.h2) + + @property + def is_derived(self) -> bool: + """True if this frequency is derived from another.""" + return self == Freq.h2 + + @property + def source_freq(self) -> Freq | None: + """The source frequency from which this freq is derived.""" + if self == Freq.h2: + return Freq.h1 + return None + + @property + def akshare_period(self) -> str: + """The period string used by akshare API.""" + mapping = { + Freq.d1: "daily", + Freq.w1: "weekly", + Freq.M1: "monthly", + } + if self in mapping: + return mapping[self] + raise ValueError(f"No akshare period for intraday freq {self.value}") + + @property + def akshare_min_period(self) -> str: + """The minute period string used by akshare stock_zh_a_hist_min_em.""" + mapping = { + Freq.m1: "1", + Freq.m5: "5", + Freq.m15: "15", + Freq.m30: "30", + Freq.h1: "60", + } + if self in mapping: + return mapping[self] + raise ValueError(f"No akshare min period for freq {self.value}") + + @property + def storage_dir(self) -> str: + """Filesystem-safe directory name (avoids case conflicts on macOS). + + '1m' stays '1m', but '1M' becomes '1mon' to distinguish from '1m'. + """ + if self == Freq.M1: + return "1mon" + return self.value + + +# Frequencies that have direct API support (not derived) +STORED_FREQS: tuple[Freq, ...] = ( + Freq.d1, Freq.w1, Freq.M1, + Freq.m1, Freq.m5, Freq.m15, Freq.m30, Freq.h1, +) + +# Frequencies that can be backfilled from akshare (full history) +BACKFILLABLE_FREQS: tuple[Freq, ...] = (Freq.d1, Freq.w1, Freq.M1) + +# Frequencies that require daily EOD accumulation (limited API history) +INTRADAY_FREQS: tuple[Freq, ...] = (Freq.m1, Freq.m5, Freq.m15, Freq.m30, Freq.h1) + +# Frequencies that are derived at query time +DERIVED_FREQS: tuple[Freq, ...] = (Freq.h2,) diff --git a/src/ashare_dp/data/__init__.py b/src/ashare_dp/data/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/ashare_dp/data/akshare_client.py b/src/ashare_dp/data/akshare_client.py new file mode 100644 index 0000000..bd02962 --- /dev/null +++ b/src/ashare_dp/data/akshare_client.py @@ -0,0 +1,321 @@ +"""AKShare wrapper with retry logic, rate limiting, and dual-backend support. + +Backends: + - em (East Money): Better data quality and history. Geo-blocked outside China. + - sina (Sina): Accessible globally. Slightly less history for minute data. + +The client auto-detects which backend is reachable and falls back from +East Money to Sina transparently. + +We clear proxy env vars AND override macOS SystemConfiguration proxy +detection so that akshare API calls always connect directly (domestic sites). +""" + +from __future__ import annotations + +import os +import re +import time +from datetime import date +from typing import Optional + +import pandas as pd +import urllib.request +from loguru import logger + +from ashare_dp.config import Settings +from ashare_dp.core.exceptions import DataSourceError + +settings = Settings() + +# === Force direct connection (no proxy) === +for _var in ("HTTP_PROXY", "HTTPS_PROXY", "http_proxy", "https_proxy", "ALL_PROXY", + "no_proxy", "NO_PROXY"): + os.environ.pop(_var, None) + +_original_getproxies = urllib.request.getproxies + + +def _getproxies_direct(): + return {} + + +urllib.request.getproxies = _getproxies_direct +logger.debug("Proxy disabled: env vars cleared + macOS SystemConfiguration bypassed") + + +def _to_akshare_date(d: date) -> str: + return d.strftime("%Y%m%d") + + +def _code_to_sina_symbol(code: str) -> str: + """Convert numeric stock code to Sina symbol format. + + '000001' -> 'sz000001', '600000' -> 'sh600000', '920000' -> 'bj920000' + """ + code = str(code).zfill(6) + if code.startswith(("4", "8")) or code.startswith("92"): + return f"bj{code}" + elif code.startswith("6") or code.startswith("9"): + return f"sh{code}" + else: + return f"sz{code}" + + +def _sina_symbol_to_ts_code(symbol: str) -> str: + """Convert Sina symbol to ts_code format. + + 'sh600000' -> '600000.SH', 'sz000001' -> '000001.SZ', 'bj920000' -> '920000.BJ' + """ + match = re.match(r'^([a-z]+)(\d{6})$', str(symbol)) + if match: + prefix, code = match.groups() + exchange = prefix.upper() + return f"{code}.{exchange}" + return str(symbol) + + +class AKShareClient: + """Wrapper around akshare with retry, error handling, and dual backend. + + Tries East Money backend first (better data), falls back to Sina. + Uses Sina for stock list and trading calendar (East Money doesn't have these). + """ + + def __init__( + self, + max_retries: int | None = None, + retry_delay: float | None = None, + backend: str | None = None, + ): + self.max_retries = max_retries or settings.akshare_max_retries + self.retry_delay = retry_delay or settings.akshare_retry_delay + self._backend = backend or settings.akshare_backend + self._em_available: bool | None = None # Cached reachability check + + # ---- Backend selection ---- + + def _resolve_backend(self) -> str: + """Determine which backend to use. + + If setting is 'auto', probes East Money and falls back to Sina. + Results are cached so the probe is only done once. + """ + if self._backend in ("em", "sina"): + return self._backend + # auto: probe East Money reachability + if self._em_available is None: + self._em_available = self._probe_em() + return "em" if self._em_available else "sina" + + def _probe_em(self) -> bool: + """Quick check if East Money is reachable.""" + try: + import requests + r = requests.get( + "https://push2.eastmoney.com/api/qt/stock/kline/get", + params={ + "secid": "1.000001", + "klt": "101", + "fqt": "1", + "beg": "20260501", + "end": "20260501", + "fields1": "f1", + "fields2": "f51", + }, + timeout=5, + ) + return r.status_code == 200 and len(r.text) > 100 + except Exception: + return False + + def _retry(self, func, *args, **kwargs): + """Execute a function with exponential backoff retry.""" + last_error = None + for attempt in range(self.max_retries + 1): + try: + return func(*args, **kwargs) + except Exception as e: + last_error = e + if attempt < self.max_retries: + wait = self.retry_delay * (2 ** attempt) + logger.warning( + f"AKShare call failed (attempt {attempt + 1}/{self.max_retries + 1}): " + f"{e}. Retrying in {wait:.0f}s..." + ) + time.sleep(wait) + else: + logger.error(f"AKShare call failed after {self.max_retries + 1} attempts: {e}") + raise DataSourceError(f"AKShare call failed: {last_error}") from last_error + + # ---- Stock Info (Sina only) ---- + + def get_stock_list(self) -> pd.DataFrame: + """Fetch all A-share stock basic info from Sina. + + Returns DataFrame with columns: code, name + """ + logger.info("Fetching A-share stock list...") + import akshare as ak + df = self._retry(lambda: ak.stock_info_a_code_name()) + logger.info(f"Fetched {len(df)} stocks") + return df + + # ---- Historical K-line (daily only; weekly/monthly derived by caller) ---- + + def get_hist( + self, + symbol: str, + period: str = "daily", + start_date: str = "19900101", + end_date: str = "20260517", + adjust: str = "qfq", + ) -> pd.DataFrame: + """Fetch historical daily K-line data for a single stock. + + Tries East Money first, falls back to Sina. + + Args: + symbol: Stock code without exchange prefix, e.g. '000001'. + period: Only 'daily' is supported. Weekly/monthly are derived. + start_date: Start date in 'YYYYMMDD' format. + end_date: End date in 'YYYYMMDD' format. + adjust: 'qfq', 'hfq', or ''. + + Returns: + DataFrame with columns: date, open, high, low, close, volume, amount + """ + backend = self._resolve_backend() + if backend == "em": + try: + return self._get_hist_em(symbol, start_date, end_date, adjust) + except Exception as e: + logger.warning(f"East Money daily failed for {symbol}, falling back to Sina: {e}") + return self._get_hist_sina(symbol, start_date, end_date, adjust) + return self._get_hist_sina(symbol, start_date, end_date, adjust) + + def _get_hist_em( + self, symbol: str, start_date: str, end_date: str, adjust: str + ) -> pd.DataFrame: + import akshare as ak + return self._retry(lambda: ak.stock_zh_a_hist( + symbol=symbol, + period="daily", + start_date=start_date, + end_date=end_date, + adjust=adjust, + )) + + def _get_hist_sina( + self, symbol: str, start_date: str, end_date: str, adjust: str + ) -> pd.DataFrame: + import akshare as ak + sina_symbol = _code_to_sina_symbol(symbol) + return self._retry(lambda: ak.stock_zh_a_daily( + symbol=sina_symbol, + start_date=start_date, + end_date=end_date, + adjust=adjust, + )) + + # ---- Minute K-line ---- + + def get_hist_min( + self, + symbol: str, + period: str = "5", + start_date: str | None = None, + end_date: str | None = None, + ) -> pd.DataFrame: + """Fetch minute-level K-line data. + + Tries East Money first (supports date range), falls back to Sina. + + Args: + symbol: Stock code without exchange prefix, e.g. '000001'. + period: One of '1', '5', '15', '30', '60'. + start_date: Start date in 'YYYYMMDD' format (East Money only). + end_date: End date in 'YYYYMMDD' format (East Money only). + + Returns: + DataFrame with minute OHLCV data. + """ + backend = self._resolve_backend() + if backend == "em": + try: + return self._get_hist_min_em(symbol, period, start_date, end_date) + except Exception as e: + logger.warning(f"East Money minute failed for {symbol}, falling back to Sina: {e}") + return self._get_hist_min_sina(symbol, period, start_date, end_date) + return self._get_hist_min_sina(symbol, period, start_date, end_date) + + def _get_hist_min_em( + self, symbol: str, period: str, + start_date: str | None, end_date: str | None, + ) -> pd.DataFrame: + import akshare as ak + kwargs = {"symbol": symbol, "period": period} + if start_date and end_date: + kwargs["start_date"] = start_date + kwargs["end_date"] = end_date + elif start_date: + kwargs["start_date"] = start_date + kwargs["end_date"] = start_date + return self._retry(lambda: ak.stock_zh_a_hist_min_em(**kwargs)) + + def _get_hist_min_sina( + self, symbol: str, period: str, + start_date: str | None, end_date: str | None, + ) -> pd.DataFrame: + import akshare as ak + sina_symbol = _code_to_sina_symbol(symbol) + df = self._retry(lambda: ak.stock_zh_a_minute( + symbol=sina_symbol, period=period, + )) + if df is None or df.empty: + return pd.DataFrame() + # Sina doesn't support date range - filter after fetch + if "day" in df.columns and (start_date or end_date): + df["day"] = pd.to_datetime(df["day"]) + if start_date: + start_ts = pd.Timestamp(start_date) + df = df[df["day"] >= start_ts] + if end_date: + end_ts = pd.Timestamp(end_date) + pd.Timedelta(days=1) + df = df[df["day"] < end_ts] + return df + + # ---- Real-time Spot ---- + + def get_spot(self) -> pd.DataFrame: + """Fetch real-time market snapshot for ALL A-share stocks. + + Tries East Money first (faster), falls back to Sina. + """ + backend = self._resolve_backend() + if backend == "em": + try: + return self._get_spot_em() + except Exception as e: + logger.warning(f"East Money spot failed, falling back to Sina: {e}") + return self._get_spot_sina() + return self._get_spot_sina() + + def _get_spot_em(self) -> pd.DataFrame: + import akshare as ak + return self._retry(lambda: ak.stock_zh_a_spot_em()) + + def _get_spot_sina(self) -> pd.DataFrame: + import akshare as ak + return self._retry(lambda: ak.stock_zh_a_spot()) + + # ---- Trading Calendar (Sina only) ---- + + def get_trading_calendar(self) -> pd.DataFrame: + """Fetch historical trading calendar from Sina. + + Returns: + DataFrame with trade_date column. + """ + import akshare as ak + return self._retry(lambda: ak.tool_trade_date_hist_sina()) diff --git a/src/ashare_dp/data/backfill.py b/src/ashare_dp/data/backfill.py new file mode 100644 index 0000000..774b567 --- /dev/null +++ b/src/ashare_dp/data/backfill.py @@ -0,0 +1,542 @@ +"""Historical backfill pipeline for A-share data. + +Pulls full history for daily K-line from Sina and stores in Parquet. +Weekly and monthly K-line are derived from stored daily data. +""" + +from __future__ import annotations + +import time +from concurrent.futures import ThreadPoolExecutor, as_completed +from datetime import date, datetime, timedelta +from typing import Optional + +import pandas as pd +from loguru import logger + +from ashare_dp.config import Settings +from ashare_dp.core.exceptions import StorageError +from ashare_dp.core.models import ( + BACKFILLABLE_FREQS, + INTRADAY_FREQS, + Freq, +) +from ashare_dp.data.akshare_client import AKShareClient +from ashare_dp.storage.database import get_db +from ashare_dp.storage.repository import KLineRepository + +settings = Settings() + +# Sina daily API returns English column names +SINA_DAILY_COLUMN_MAPPING = { + "date": "trade_time", + "open": "open", + "high": "high", + "low": "low", + "close": "close", + "volume": "volume", + "amount": "amount", +} + +# Sina minute API returns these columns +SINA_MIN_COLUMN_MAPPING = { + "day": "trade_time", + "open": "open", + "high": "high", + "low": "low", + "close": "close", + "volume": "volume", + "amount": "amount", +} + +# Legacy Chinese column mappings (for backward compatibility) +LEGACY_COLUMN_MAPPING = { + "日期": "trade_time", + "股票代码": "ts_code", + "开盘": "open", + "最高": "high", + "最低": "low", + "收盘": "close", + "成交量": "volume", + "成交额": "amount", +} + +LEGACY_MIN_COLUMN_MAPPING = { + "时间": "trade_time", + "开盘": "open", + "最高": "high", + "最低": "low", + "收盘": "close", + "成交量": "volume", + "成交额": "amount", +} + +REQUIRED_COLS = ["ts_code", "trade_time", "open", "high", "low", "close", "volume", "amount"] + + +def _normalize_hist_df(df: pd.DataFrame, symbol: str, freq: Freq) -> pd.DataFrame: + """Normalize daily K-line output to standard K-line schema.""" + df = df.copy() + + # Try English mapping first (Sina), then Chinese (legacy) + for mapping in [SINA_DAILY_COLUMN_MAPPING, LEGACY_COLUMN_MAPPING]: + rename_map = {} + for src, dst in mapping.items(): + if src in df.columns: + rename_map[src] = dst + if rename_map: + df = df.rename(columns=rename_map) + break + + if "ts_code" not in df.columns: + df["ts_code"] = symbol + + if "trade_time" in df.columns: + df["trade_time"] = pd.to_datetime(df["trade_time"]) + df["trade_date"] = df["trade_time"].dt.date + elif "trade_date" in df.columns: + df["trade_date"] = pd.to_datetime(df["trade_date"]).dt.date + df["trade_time"] = pd.to_datetime(df["trade_date"]) + + # Ensure numeric columns + for col in ["open", "high", "low", "close", "volume", "amount"]: + if col in df.columns: + df[col] = pd.to_numeric(df[col], errors="coerce") + + return df + + +def _normalize_min_df(df: pd.DataFrame, symbol: str) -> pd.DataFrame: + """Normalize minute K-line output to standard K-line schema.""" + df = df.copy() + + # Try English mapping first (Sina), then Chinese (legacy) + for mapping in [SINA_MIN_COLUMN_MAPPING, LEGACY_MIN_COLUMN_MAPPING]: + rename_map = {} + for src, dst in mapping.items(): + if src in df.columns: + rename_map[src] = dst + if rename_map: + df = df.rename(columns=rename_map) + break + + if "ts_code" not in df.columns: + df["ts_code"] = symbol + + if "trade_time" in df.columns: + df["trade_time"] = pd.to_datetime(df["trade_time"]) + df["trade_date"] = df["trade_time"].dt.date + + for col in ["open", "high", "low", "close", "volume", "amount"]: + if col in df.columns: + df[col] = pd.to_numeric(df[col], errors="coerce") + + return df + + +def _resample_daily_to_period(df: pd.DataFrame, freq: Freq) -> pd.DataFrame: + """Resample daily K-line data to weekly or monthly. + + Args: + df: DataFrame with standard K-line columns. + freq: Target frequency (w1 or M1). + + Returns: + Resampled DataFrame. + """ + if df.empty: + return df + + df = df.copy() + df["trade_time"] = pd.to_datetime(df["trade_time"]) + df = df.set_index("trade_time") + + if freq == Freq.w1: + group_key = pd.Grouper(freq="W") + elif freq == Freq.M1: + group_key = pd.Grouper(freq="ME") + else: + raise ValueError(f"Unsupported resample frequency: {freq}") + + grouped = df.groupby(["ts_code", group_key]) + + result = grouped.agg({ + "open": "first", + "high": "max", + "low": "min", + "close": "last", + "volume": "sum", + "amount": "sum", + }).reset_index() + + result["trade_date"] = result["trade_time"].dt.date + result["freq"] = freq.value + + return result + + +class BackfillPipeline: + """Orchestrates historical data backfill.""" + + def __init__( + self, + client: AKShareClient | None = None, + repo: KLineRepository | None = None, + max_workers: int | None = None, + ): + self.client = client or AKShareClient() + self.repo = repo or KLineRepository() + self.max_workers = max_workers or settings.backfill_workers + self._db = get_db() + + def init_db(self): + """Initialize database schema (tables and views).""" + from ashare_dp.storage.schema import DDL_STATEMENTS + + logger.info("Initializing database schema...") + for ddl in DDL_STATEMENTS: + try: + self._db.execute(ddl) + except Exception as e: + logger.warning(f"DDL warning: {e}") + logger.info("Database schema initialized") + + def load_stock_list(self) -> pd.DataFrame: + """Fetch stock list and store in DuckDB.""" + logger.info("Loading stock list...") + df = self.client.get_stock_list() + + if df.empty: + logger.warning("No stocks returned from akshare") + return df + + # Normalize columns + df = df.rename(columns={ + "code": "symbol", + "name": "name", + }) + + # Build ts_code from code + def _make_ts_code(code: str) -> str: + code = str(code).zfill(6) + if code.startswith(("4", "8")) or code.startswith("92"): + return f"{code}.BJ" + elif code.startswith("6") or code.startswith("9"): + return f"{code}.SH" + else: + return f"{code}.SZ" + + if "symbol" in df.columns: + df["ts_code"] = df["symbol"].apply(_make_ts_code) + df["exchange"] = df["ts_code"].str[-2:] + + # Upsert into DuckDB + if "ts_code" in df.columns and "symbol" in df.columns: + for _, row in df.iterrows(): + try: + self._db.execute(""" + INSERT OR REPLACE INTO stock_info (ts_code, symbol, name, exchange, updated_at) + VALUES (?, ?, ?, ?, CURRENT_TIMESTAMP) + """, ( + row.get("ts_code"), + str(row.get("symbol", "")), + str(row.get("name", "")), + row.get("exchange", ""), + )) + except Exception: + pass + + logger.info(f"Loaded {len(df)} stocks into stock_info") + return df + + def load_trading_calendar(self): + """Fetch trading calendar and store in DuckDB.""" + logger.info("Loading trading calendar...") + df = self.client.get_trading_calendar() + + if df.empty: + logger.warning("Empty trading calendar returned") + return + + df["trade_date"] = pd.to_datetime(df["trade_date"]).dt.date + + for _, row in df.iterrows(): + d = row["trade_date"] + try: + self._db.execute(""" + INSERT OR REPLACE INTO trading_calendar (trade_date, is_trading_day, week_day, year, month) + VALUES (?, 1, ?, ?, ?) + """, ( + d, + d.weekday(), + d.year, + d.month, + )) + except Exception: + pass + + logger.info(f"Loaded {len(df)} trading days into trading_calendar") + + def backfill_daily_weekly_monthly( + self, + symbols: list[str] | None = None, + start_date: date | None = None, + end_date: date | None = None, + ) -> dict: + """Backfill daily K-line, then derive weekly and monthly. + + Daily data is fetched from Sina API. Weekly and monthly are + derived by resampling the daily data. + + Args: + symbols: List of stock symbols (e.g. ['000001', '600000']). + If None, backfills all stocks from stock_info. + start_date: Start date for backfill (default: 1990-01-01). + end_date: End date for backfill (default: today). + + Returns: + Dict with summary stats per frequency. + """ + if symbols is None: + stocks = self.client.get_stock_list() + symbols = [str(c).zfill(6) for c in stocks["code"].tolist()] + + if start_date is None: + start_date = date(1990, 1, 1) + if end_date is None: + end_date = date.today() + + start_str = start_date.strftime("%Y%m%d") + end_str = end_date.strftime("%Y%m%d") + + logger.info( + f"Starting backfill: {len(symbols)} stocks, " + f"{start_str} to {end_str}, {self.max_workers} workers" + ) + + results = {} + + # ---- Step 1: Backfill daily K-line via Sina API ---- + logger.info("Backfilling 1d (daily) from Sina...") + completed = 0 + failed = 0 + all_daily_frames = [] + + def _backfill_daily(symbol: str): + try: + df = self.client.get_hist( + symbol=symbol, + period="daily", + start_date=start_str, + end_date=end_str, + adjust="qfq", + ) + if df is not None and not df.empty: + return _normalize_hist_df(df, symbol, Freq.d1) + return None + except Exception as e: + logger.error(f"Failed backfill {symbol} 1d: {e}") + raise + + with ThreadPoolExecutor(max_workers=self.max_workers) as executor: + futures = { + executor.submit(_backfill_daily, sym): sym + for sym in symbols + } + for future in as_completed(futures): + sym = futures[future] + try: + df = future.result() + if df is not None and not df.empty: + all_daily_frames.append(df) + completed += 1 + except Exception: + failed += 1 + + if (completed + failed) % 100 == 0: + logger.info( + f" 1d: {completed + failed}/{len(symbols)} " + f"({completed} ok, {failed} fail)" + ) + + # Batch write all daily data at once + total_records = 0 + if all_daily_frames: + combined = pd.concat(all_daily_frames, ignore_index=True) + total_records = self.repo.write_klines(combined, Freq.d1) + + results[Freq.d1.value] = { + "records": total_records, + "completed": completed, + "failed": failed, + } + logger.info( + f" 1d done: {total_records} records, " + f"{completed} stocks ok, {failed} failed" + ) + + # ---- Step 2: Derive weekly from stored daily data ---- + logger.info("Deriving 1w (weekly) from daily data...") + self._derive_from_daily(Freq.w1, symbols, start_str, end_str, results) + + # ---- Step 3: Derive monthly from stored daily data ---- + logger.info("Deriving 1M (monthly) from daily data...") + self._derive_from_daily(Freq.M1, symbols, start_str, end_str, results) + + return results + + def _derive_from_daily( + self, + freq: Freq, + symbols: list[str], + start_str: str, + end_str: str, + results: dict, + ): + """Derive weekly/monthly K-lines from stored daily data. + + Reads daily data for all symbols, resamples, and writes in batch + to avoid file overwrite issues with the grouped-per-day write strategy. + """ + total_records = 0 + completed = 0 + failed = 0 + all_frames = [] + + for symbol in symbols: + try: + df = self.repo.read_klines( + freq=Freq.d1, + ts_code=symbol, + start_date=date.fromisoformat( + pd.Timestamp(start_str).strftime("%Y-%m-%d") + ) if len(start_str) == 8 else date.fromisoformat(start_str), + end_date=date.fromisoformat( + pd.Timestamp(end_str).strftime("%Y-%m-%d") + ) if len(end_str) == 8 else date.fromisoformat(end_str), + limit=1_000_000, + ) + if df is not None and not df.empty: + df = _resample_daily_to_period(df, freq) + all_frames.append(df) + completed += 1 + except Exception as e: + logger.error(f"Failed derive {freq.value} for {symbol}: {e}") + failed += 1 + + if (completed + failed) % 500 == 0: + logger.info( + f" {freq.value}: {completed + failed}/{len(symbols)} " + f"({completed} ok, {failed} fail)" + ) + + # Batch write all derived data at once + if all_frames: + combined = pd.concat(all_frames, ignore_index=True) + total_records = self.repo.write_klines(combined, freq) + + results[freq.value] = { + "records": total_records, + "completed": completed, + "failed": failed, + } + logger.info( + f" {freq.value} done: {total_records} records, " + f"{completed} stocks ok, {failed} failed" + ) + + def backfill_minute( + self, + symbols: list[str] | None = None, + days_back: int = 30, + ) -> dict: + """Backfill recent minute data for all (or specified) stocks. + + Uses Sina minute API which returns all recent data. + We filter to the requested date range after fetching. + + Args: + symbols: List of stock symbols. If None, backfills all. + days_back: Number of calendar days to look back. + + Returns: + Dict with summary stats per frequency. + """ + if symbols is None: + stocks = self.client.get_stock_list() + symbols = [str(c).zfill(6) for c in stocks["code"].tolist()] + + end_date = date.today() + lookback = end_date - timedelta(days=days_back) + + start_str = lookback.strftime("%Y%m%d") + end_str = end_date.strftime("%Y%m%d") + + logger.info( + f"Starting minute backfill: {len(symbols)} stocks, " + f"{start_str} to {end_str}, {self.max_workers} workers" + ) + + results = {} + + for freq in INTRADAY_FREQS: + period = freq.akshare_min_period + logger.info(f"Backfilling {freq.value} ({period}min)...") + + completed = 0 + failed = 0 + all_min_frames = [] + + def _backfill_min(symbol: str): + try: + df = self.client.get_hist_min( + symbol=symbol, + period=period, + start_date=start_str, + end_date=end_str, + ) + if df is not None and not df.empty: + return _normalize_min_df(df, symbol) + return None + except Exception as e: + logger.error(f"Failed minute backfill {symbol} {freq.value}: {e}") + raise + + with ThreadPoolExecutor(max_workers=self.max_workers) as executor: + futures = { + executor.submit(_backfill_min, sym): sym + for sym in symbols + } + for future in as_completed(futures): + sym = futures[future] + try: + df = future.result() + if df is not None and not df.empty: + all_min_frames.append(df) + completed += 1 + except Exception: + failed += 1 + + if (completed + failed) % 100 == 0: + logger.info( + f" {freq.value}: {completed + failed}/{len(symbols)} " + f"({completed} ok, {failed} fail)" + ) + + # Batch write all data for this frequency at once + total_records = 0 + if all_min_frames: + combined = pd.concat(all_min_frames, ignore_index=True) + total_records = self.repo.write_klines(combined, freq) + + results[freq.value] = { + "records": total_records, + "completed": completed, + "failed": failed, + } + logger.info( + f" {freq.value} done: {total_records} records, " + f"{completed} stocks ok, {failed} failed" + ) + + return results diff --git a/src/ashare_dp/data/eod.py b/src/ashare_dp/data/eod.py new file mode 100644 index 0000000..b837718 --- /dev/null +++ b/src/ashare_dp/data/eod.py @@ -0,0 +1,176 @@ +"""End-of-day batch data pull. + +Pulls daily and minute K-line data for ALL active stocks after market close. +Weekly and monthly are derived from stored daily data. +All writes are batched to avoid per-file overwrites. +""" + +from __future__ import annotations + +from datetime import date, datetime, timedelta + +import pandas as pd +from loguru import logger + +from ashare_dp.core.calendar import BEIJING_TZ +from ashare_dp.core.models import BACKFILLABLE_FREQS, INTRADAY_FREQS, Freq +from ashare_dp.data.akshare_client import AKShareClient +from ashare_dp.data.backfill import _normalize_hist_df, _normalize_min_df, _resample_daily_to_period +from ashare_dp.storage.repository import KLineRepository + + +class EODPipeline: + """Pulls end-of-day data for all active stocks.""" + + def __init__( + self, + client: AKShareClient | None = None, + repo: KLineRepository | None = None, + ): + self.client = client or AKShareClient() + self.repo = repo or KLineRepository() + + def get_active_symbols(self) -> list[str]: + """Get list of currently listed stock symbols.""" + stocks = self.client.get_stock_list() + return [str(c).zfill(6) for c in stocks["code"].tolist()] + + def pull_daily(self, symbols: list[str], trade_date: date) -> dict: + """Pull daily K-line for all stocks for a specific trading day. + + Collects all stocks' data then writes in one batch. + """ + date_str = trade_date.strftime("%Y%m%d") + logger.info(f"EOD daily pull: {len(symbols)} stocks for {date_str}") + + frames = [] + completed = 0 + failed = 0 + + for symbol in symbols: + try: + df = self.client.get_hist( + symbol=symbol, + period="daily", + start_date=date_str, + end_date=date_str, + adjust="qfq", + ) + if df is not None and not df.empty: + frames.append(_normalize_hist_df(df, symbol, Freq.d1)) + completed += 1 + except Exception as e: + logger.error(f"EOD daily failed for {symbol}: {e}") + failed += 1 + + total_records = 0 + if frames: + combined = pd.concat(frames, ignore_index=True) + total_records = self.repo.write_klines(combined, Freq.d1, partition_date=trade_date) + + logger.info(f"EOD daily done: {total_records} records, {completed} ok, {failed} failed") + return {"records": total_records, "completed": completed, "failed": failed} + + def pull_minute(self, symbols: list[str], trade_date: date) -> dict: + """Pull minute K-line for all stocks for a specific trading day. + + Pulls all intraday frequencies: 1m, 5m, 15m, 30m, 1h. + Collects all stocks' data then batch-writes per frequency. + """ + date_str = trade_date.strftime("%Y%m%d") + logger.info(f"EOD minute pull: {len(symbols)} stocks for {date_str}") + + results = {} + + for freq in INTRADAY_FREQS: + period = freq.akshare_min_period + frames = [] + completed = 0 + failed = 0 + + for symbol in symbols: + try: + df = self.client.get_hist_min( + symbol=symbol, + period=period, + start_date=date_str, + end_date=date_str, + ) + if df is not None and not df.empty: + frames.append(_normalize_min_df(df, symbol)) + completed += 1 + except Exception as e: + logger.error(f"EOD {freq.value} failed for {symbol}: {e}") + failed += 1 + + total_records = 0 + if frames: + combined = pd.concat(frames, ignore_index=True) + total_records = self.repo.write_klines(combined, freq, partition_date=trade_date) + + results[freq.value] = { + "records": total_records, + "completed": completed, + "failed": failed, + } + logger.info( + f"EOD {freq.value} done: {total_records} records, " + f"{completed} ok, {failed} failed" + ) + + return results + + def pull_weekly_monthly(self, symbols: list[str], trade_date: date) -> dict: + """Derive weekly and monthly K-line from stored daily data. + + Reads recent daily data for all symbols, resamples, and writes in batch. + Only keeps the latest bar for each stock. + """ + date_str = trade_date.strftime("%Y%m%d") + logger.info(f"EOD weekly/monthly derive from daily for {date_str}") + + results = {} + + for freq in [Freq.w1, Freq.M1]: + frames = [] + completed = 0 + failed = 0 + + lookback = 60 if freq == Freq.w1 else 90 + start_d = trade_date - timedelta(days=lookback) + + for symbol in symbols: + try: + df = self.repo.read_klines( + freq=Freq.d1, + ts_code=symbol, + start_date=start_d, + end_date=trade_date, + limit=500, + ) + if df is not None and not df.empty: + df = _resample_daily_to_period(df, freq) + if not df.empty: + # Only keep the latest bar per stock + frames.append(df.iloc[-1:]) + completed += 1 + except Exception as e: + logger.error(f"EOD {freq.value} derive failed for {symbol}: {e}") + failed += 1 + + total_records = 0 + if frames: + combined = pd.concat(frames, ignore_index=True) + total_records = self.repo.write_klines(combined, freq, partition_date=trade_date) + + results[freq.value] = { + "records": total_records, + "completed": completed, + "failed": failed, + } + logger.info( + f"EOD {freq.value} done: {total_records} records, " + f"{completed} ok, {failed} failed" + ) + + return results diff --git a/src/ashare_dp/data/realtime.py b/src/ashare_dp/data/realtime.py new file mode 100644 index 0000000..a542a72 --- /dev/null +++ b/src/ashare_dp/data/realtime.py @@ -0,0 +1,148 @@ +"""Real-time spot data polling service. + +Background asyncio task that polls akshare spot data and broadcasts +to subscribed WebSocket clients during trading hours. +""" + +from __future__ import annotations + +import asyncio +from datetime import datetime + +from loguru import logger + +from ashare_dp.api.websocket.manager import manager +from ashare_dp.core.calendar import BEIJING_TZ, MarketState, determine_market_state +from ashare_dp.data.akshare_client import AKShareClient + + +def _code_to_ts_code(code: str) -> str: + """Convert a raw stock code to ts_code format. + + Handles both Sina format (with exchange prefix, e.g. 'sh600000') + and numeric-only format (e.g. '000001'). + """ + code = str(code) + # Sina format: 'sh600000', 'sz000001', 'bj920000' + if code[:2].isalpha() and len(code) == 8: + prefix = code[:2].lower() + num = code[2:] + exchange_map = {"sh": "SH", "sz": "SZ", "bj": "BJ"} + return f"{num}.{exchange_map.get(prefix, prefix.upper())}" + # Numeric format + code = code.zfill(6) + if code.startswith(("4", "8")) or code.startswith("92"): + return f"{code}.BJ" + elif code.startswith("6") or code.startswith("9"): + return f"{code}.SH" + return f"{code}.SZ" + + +class RealtimePoller: + """Background task that polls market data and broadcasts to subscribers.""" + + def __init__( + self, + client: AKShareClient | None = None, + poll_interval: int = 5, + ): + self._client = client or AKShareClient() + self._poll_interval = poll_interval + self._running = False + self._task: asyncio.Task | None = None + + async def start(self): + """Start the polling loop.""" + self._running = True + self._task = asyncio.create_task(self._poll_loop()) + logger.info("RealtimePoller started") + + async def stop(self): + """Stop the polling loop.""" + self._running = False + if self._task: + self._task.cancel() + try: + await self._task + except asyncio.CancelledError: + pass + logger.info("RealtimePoller stopped") + + def _get_poll_interval(self, state: str) -> int: + """Get polling interval based on market state.""" + intervals = { + MarketState.TRADING: self._poll_interval, + MarketState.PRE_OPEN: 3, + MarketState.CLOSING_AUCTION: 3, + MarketState.LUNCH_BREAK: 60, + MarketState.CLOSED: 60, + } + return intervals.get(state, 60) + + async def _poll_loop(self): + """Main polling loop.""" + last_heartbeat = datetime.now(BEIJING_TZ) + + while self._running: + try: + now = datetime.now(BEIJING_TZ) + state = determine_market_state(now) + + # Broadcast market state to all clients + if manager.active_connections > 0: + await manager.broadcast({ + "type": "market_state", + "state": state, + "timestamp": now.isoformat(), + }) + + # Poll spot data during trading periods + if state in (MarketState.TRADING, MarketState.PRE_OPEN, MarketState.CLOSING_AUCTION): + codes = await manager.get_all_subscribed_codes() + if codes: + try: + df = await asyncio.to_thread(self._client.get_spot) + if df is not None and not df.empty: + # Normalize and filter + if "代码" in df.columns: + df["ts_code"] = df["代码"].apply(_code_to_ts_code) + + spot_data = {} + for _, row in df.iterrows(): + tc = row.get("ts_code", "") + if tc in codes: + spot_data[tc] = { + "price": float(row.get("最新价", 0) or 0), + "change": float(row.get("涨跌额", 0) or 0), + "pct_chg": float(row.get("涨跌幅", 0) or 0), + "volume": int(row.get("成交量", 0) or 0), + "amount": float(row.get("成交额", 0) or 0), + "high": float(row.get("最高", 0) or 0), + "low": float(row.get("最低", 0) or 0), + "open": float(row.get("今开", 0) or 0), + "pre_close": float(row.get("昨收", 0) or 0), + "name": str(row.get("名称", "")), + } + + if spot_data: + await manager.broadcast_filtered(spot_data) + except Exception as e: + logger.error(f"Spot poll error: {e}") + + # Heartbeat every 30 seconds + if (now - last_heartbeat).total_seconds() >= 30: + await manager.broadcast({ + "type": "heartbeat", + "timestamp": now.isoformat(), + "connections": manager.active_connections, + }) + last_heartbeat = now + + except Exception as e: + logger.error(f"Poll loop error: {e}") + + await asyncio.sleep(self._get_poll_interval(state)) + + +# Global poller instance +poller = RealtimePoller() diff --git a/src/ashare_dp/scheduler/__init__.py b/src/ashare_dp/scheduler/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/ashare_dp/scheduler/jobs.py b/src/ashare_dp/scheduler/jobs.py new file mode 100644 index 0000000..5fbabd7 --- /dev/null +++ b/src/ashare_dp/scheduler/jobs.py @@ -0,0 +1,77 @@ +"""Scheduler job implementations.""" + +from __future__ import annotations + +from datetime import date, datetime + +from loguru import logger + +from ashare_dp.core.calendar import BEIJING_TZ +from ashare_dp.data.eod import EODPipeline +from ashare_dp.storage.repository import KLineRepository + + +async def eod_pull_job(): + """End-of-day data pull job. + + Triggered at 15:05 Beijing time on trading days. + Pulls daily, minute, weekly, and monthly data for all active stocks. + """ + now = datetime.now(BEIJING_TZ) + today = now.date() + + # Verify today is a weekday (trading day check done by scheduler) + if now.weekday() >= 5: + logger.info("EOD: Weekend, skipping") + return + + logger.info(f"EOD job starting for {today.isoformat()}...") + + pipeline = EODPipeline() + + try: + symbols = pipeline.get_active_symbols() + logger.info(f"EOD: {len(symbols)} active stocks") + except Exception as e: + logger.error(f"EOD: Failed to get stock list: {e}") + return + + # Pull daily data + try: + pipeline.pull_daily(symbols, today) + except Exception as e: + logger.error(f"EOD daily pull failed: {e}") + + # Pull minute data (1m, 5m, 15m, 30m, 1h) + try: + pipeline.pull_minute(symbols, today) + except Exception as e: + logger.error(f"EOD minute pull failed: {e}") + + # Pull weekly/monthly + try: + pipeline.pull_weekly_monthly(symbols, today) + except Exception as e: + logger.error(f"EOD weekly/monthly pull failed: {e}") + + logger.info(f"EOD job completed for {today.isoformat()}") + + +async def health_check_job(): + """Daily health check: report database statistics.""" + logger.info("Health check running...") + repo = KLineRepository() + try: + for freq_res in [repo.get_date_range(f) for f in [Freq.d1, Freq.h1, Freq.m5]]: + pass + logger.info( + f"Health check: DB OK, " + f"1d records={repo.count_records(Freq.d1)}, " + f"1h records={repo.count_records(Freq.h1)}" + ) + except Exception as e: + logger.error(f"Health check failed: {e}") + + +# Import at bottom to avoid circular +from ashare_dp.core.models import Freq diff --git a/src/ashare_dp/scheduler/scheduler.py b/src/ashare_dp/scheduler/scheduler.py new file mode 100644 index 0000000..4333570 --- /dev/null +++ b/src/ashare_dp/scheduler/scheduler.py @@ -0,0 +1,62 @@ +"""APScheduler setup for EOD and health check jobs.""" + +from __future__ import annotations + +from apscheduler.schedulers.asyncio import AsyncIOScheduler +from apscheduler.triggers.cron import CronTrigger +from loguru import logger + +from ashare_dp.core.calendar import BEIJING_TZ +from ashare_dp.data.akshare_client import AKShareClient + + +class Scheduler: + """Manages scheduled jobs using APScheduler.""" + + def __init__(self, client: AKShareClient | None = None): + self._scheduler = AsyncIOScheduler(timezone=BEIJING_TZ) + self._client = client or AKShareClient() + + def start(self): + """Start the scheduler and register jobs.""" + from ashare_dp.scheduler.jobs import eod_pull_job, health_check_job + + # EOD job: 15:05 Beijing time, Mon-Fri + self._scheduler.add_job( + eod_pull_job, + trigger=CronTrigger( + day_of_week="mon-fri", + hour=15, + minute=5, + timezone=BEIJING_TZ, + ), + id="eod_pull", + name="End-of-day data pull", + replace_existing=True, + ) + + # Health check: 8:00 Beijing time daily + self._scheduler.add_job( + health_check_job, + trigger=CronTrigger( + hour=8, + minute=0, + timezone=BEIJING_TZ, + ), + id="health_check", + name="Daily health check", + replace_existing=True, + ) + + self._scheduler.start() + logger.info("Scheduler started with EOD (15:05 Mon-Fri) + health check (08:00 daily)") + + def shutdown(self): + """Shut down the scheduler.""" + if self._scheduler.running: + self._scheduler.shutdown(wait=False) + logger.info("Scheduler shut down") + + @property + def running(self) -> bool: + return self._scheduler.running diff --git a/src/ashare_dp/storage/__init__.py b/src/ashare_dp/storage/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/ashare_dp/storage/database.py b/src/ashare_dp/storage/database.py new file mode 100644 index 0000000..93d0e9f --- /dev/null +++ b/src/ashare_dp/storage/database.py @@ -0,0 +1,98 @@ +"""DuckDB connection management. + +Uses a module-level connection. DuckDB connections are not thread-safe, +so all operations are serialized through a threading lock for writes. +""" + +from __future__ import annotations + +import threading +from pathlib import Path +from typing import Optional + +import duckdb +from loguru import logger + +from ashare_dp.config import Settings + +settings = Settings() + +# Module-level lock for write serialization +_write_lock = threading.Lock() + + +class Database: + """Manages a persistent DuckDB connection.""" + + def __init__(self, db_path: str | None = None): + self.db_path = str(db_path or settings.duckdb_path) + self._conn: Optional[duckdb.DuckDBPyConnection] = None + + @property + def conn(self) -> duckdb.DuckDBPyConnection: + """Get the current connection, creating it if needed.""" + if self._conn is None: + self.connect() + return self._conn + + def connect(self) -> duckdb.DuckDBPyConnection: + """Open a persistent connection to the DuckDB database.""" + Path(self.db_path).parent.mkdir(parents=True, exist_ok=True) + self._conn = duckdb.connect(self.db_path) + logger.info(f"Connected to DuckDB: {self.db_path}") + return self._conn + + def close(self): + """Close the database connection.""" + if self._conn is not None: + self._conn.close() + self._conn = None + logger.info("DuckDB connection closed") + + def execute(self, sql: str, params: tuple | None = None): + """Execute a SQL statement.""" + if params: + return self.conn.execute(sql, params) + return self.conn.execute(sql) + + def query(self, sql: str, params: tuple | None = None) -> list[tuple]: + """Execute a query and return all rows.""" + if params: + return self.conn.execute(sql, params).fetchall() + return self.conn.execute(sql).fetchall() + + def query_df(self, sql: str) -> "duckdb.DuckDBPyRelation": + """Execute a query and return a DuckDB relation (can convert to df).""" + return self.conn.sql(sql) + + def table_exists(self, name: str) -> bool: + """Check if a table exists.""" + result = self.conn.execute( + "SELECT count(*) FROM information_schema.tables " + "WHERE table_name = ?", + (name,) + ).fetchone() + return result[0] > 0 + + def write_lock(self): + """Acquire the write lock (use as context manager).""" + return _write_lock + + def __enter__(self): + return self + + def __exit__(self, *args): + self.close() + + +# Global database instance +_db: Optional[Database] = None + + +def get_db() -> Database: + """Get or create the global database instance.""" + global _db + if _db is None: + _db = Database() + _db.connect() + return _db diff --git a/src/ashare_dp/storage/partitioning.py b/src/ashare_dp/storage/partitioning.py new file mode 100644 index 0000000..9a58f8c --- /dev/null +++ b/src/ashare_dp/storage/partitioning.py @@ -0,0 +1,58 @@ +"""Parquet partition path management. + +Handles Hive-style partitioning paths: kline_{freq}/year=YYYY/month=MM/day=DD/ +""" + +from __future__ import annotations + +from datetime import date +from pathlib import Path + +from ashare_dp.config import Settings +from ashare_dp.core.models import Freq + +settings = Settings() + + +def partition_path(freq: Freq, d: date) -> Path: + """Build the Hive-partitioned directory path for a given frequency and date. + + Returns: data/parquet/kline_{freq}/year=YYYY/month=MM/day=DD/ + """ + return ( + Path(settings.parquet_dir) + / f"kline_{freq.storage_dir}" + / f"year={d.year}" + / f"month={d.month:02d}" + / f"day={d.day:02d}" + ) + + +def partition_glob(freq: Freq) -> str: + """Glob pattern to match all Parquet files for a given frequency. + + Returns: data/parquet/kline_{freq}/**/*.parquet + """ + return f"{settings.parquet_dir}/kline_{freq.storage_dir}/**/*.parquet" + + +def ensure_partition_dir(freq: Freq, d: date) -> Path: + """Create the partition directory if it doesn't exist and return the path.""" + p = partition_path(freq, d) + p.mkdir(parents=True, exist_ok=True) + return p + + +def parse_partition_from_path(path: str) -> dict: + """Parse Hive partition keys from a file path. + + E.g., 'data/parquet/kline_1d/year=2025/month=01/day=15/data.parquet' + returns {'year': 2025, 'month': 1, 'day': 15} + """ + parts = Path(path).parts + result = {} + for part in parts: + if "=" in part: + key, val = part.split("=", 1) + result[key] = int(val) + return result diff --git a/src/ashare_dp/storage/repository.py b/src/ashare_dp/storage/repository.py new file mode 100644 index 0000000..53e481a --- /dev/null +++ b/src/ashare_dp/storage/repository.py @@ -0,0 +1,271 @@ +"""High-level data access: write/read K-line data, stock info CRUD.""" + +from __future__ import annotations + +from datetime import date, datetime +from pathlib import Path +from typing import Optional + +import pandas as pd +import pyarrow as pa +import pyarrow.parquet as pq +from loguru import logger + +from ashare_dp.config import Settings +from ashare_dp.core.exceptions import QueryError, StorageError +from ashare_dp.core.models import DERIVED_FREQS, Freq +from ashare_dp.storage.database import get_db +from ashare_dp.storage.partitioning import ( + ensure_partition_dir, + partition_glob, +) + +settings = Settings() + +# Parquet write options: Zstd compression +PARQUET_WRITE_OPTIONS = { + "compression": "zstd", + "compression_level": 3, + "row_group_size": 100_000, +} + +# Standard K-line columns we always keep +STANDARD_COLS = [ + "ts_code", "trade_time", "trade_date", "open", "high", "low", + "close", "volume", "amount", "freq", +] + + +def _standardize_df(df: pd.DataFrame) -> pd.DataFrame: + """Drop non-standard columns to avoid schema conflicts.""" + cols = [c for c in STANDARD_COLS if c in df.columns] + return df[cols].copy() + + +class KLineRepository: + """Read/write K-line data from/to Parquet files.""" + + def __init__(self): + self._db = get_db() + + # ---- Write ---- + + def write_klines( + self, + df: pd.DataFrame, + freq: Freq, + partition_date: date | None = None, + ) -> int: + """Write K-line records to Hive-partitioned Parquet. + + All stocks for a given day are written to a single Parquet file. + Schema is standardized via _standardize_df to ensure consistency. + + Args: + df: DataFrame with K-line data. + freq: The K-line frequency. + partition_date: Date to partition by (defaults to trade_date). + + Returns: + Number of records written. + """ + if freq in DERIVED_FREQS: + raise StorageError(f"Cannot write derived frequency {freq.value}") + + if df.empty: + return 0 + + # Standardize: add freq, keep only standard columns + df = df.copy() + df["freq"] = freq.value + df = _standardize_df(df) + + # Determine partition date + if "trade_date" not in df.columns and partition_date is None: + raise StorageError("DataFrame must have 'trade_date' column or partition_date must be provided") + + records_written = 0 + + if partition_date is not None: + # Single date: all records go to one file + partition_dir = ensure_partition_dir(freq, partition_date) + file_path = partition_dir / "data.parquet" + table = pa.Table.from_pandas(df, preserve_index=False) + pq.write_table(table, str(file_path), **PARQUET_WRITE_OPTIONS) + records_written = len(df) + else: + # Group by trade_date, one file per day + df["_pd"] = pd.to_datetime(df["trade_date"]).dt.date + for d, group in df.groupby("_pd"): + group = group.drop(columns=["_pd"]) + partition_dir = ensure_partition_dir(freq, d) + file_path = partition_dir / "data.parquet" + table = pa.Table.from_pandas(group, preserve_index=False) + pq.write_table(table, str(file_path), **PARQUET_WRITE_OPTIONS) + records_written += len(group) + + logger.debug(f"Wrote {records_written} records to kline_{freq.storage_dir}") + return records_written + + # ---- Read ---- + + def read_klines( + self, + freq: Freq, + ts_code: str | None = None, + start_date: date | None = None, + end_date: date | None = None, + limit: int = 10000, + offset: int = 0, + ) -> pd.DataFrame: + """Read K-line data from Parquet files.""" + if freq == Freq.h2: + return self._read_2h(ts_code, start_date, end_date, limit, offset) + + glob = partition_glob(freq) + parquet_path = Path(settings.parquet_dir) / f"kline_{freq.storage_dir}" + if not parquet_path.exists(): + return pd.DataFrame(columns=STANDARD_COLS) + + conditions = [] + params = [] + if ts_code: + conditions.append(f"ts_code = ${len(params) + 1}") + params.append(ts_code) + if start_date: + conditions.append(f"trade_date >= ${len(params) + 1}") + params.append(start_date.isoformat()) + if end_date: + conditions.append(f"trade_date <= ${len(params) + 1}") + params.append(end_date.isoformat()) + + where_clause = "" + if conditions: + where_clause = "WHERE " + " AND ".join(conditions) + + sql = f""" + SELECT * FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true) + {where_clause} + ORDER BY trade_time + LIMIT {limit} OFFSET {offset} + """ + try: + return self._db.conn.execute(sql, params).fetchdf() + except Exception as e: + logger.warning(f"Query failed for freq={freq.value}: {e}") + return pd.DataFrame(columns=STANDARD_COLS) + + # ---- 2h Derived ---- + + def _read_2h( + self, + ts_code: str | None = None, + start_date: date | None = None, + end_date: date | None = None, + limit: int = 10000, + offset: int = 0, + ) -> pd.DataFrame: + """Derive 2h K-line from 1h data.""" + df_1h = self.read_klines( + freq=Freq.h1, + ts_code=ts_code, + start_date=start_date, + end_date=end_date, + limit=limit * 10, + offset=0, + ) + + if df_1h.empty: + return pd.DataFrame(columns=STANDARD_COLS) + + df_1h["trade_time"] = pd.to_datetime(df_1h["trade_time"]) + df_1h = df_1h.set_index("trade_time") + grouped = df_1h.groupby(["ts_code", pd.Grouper(freq="2h", level="trade_time")]) + + result = grouped.agg({ + "open": "first", + "high": "max", + "low": "min", + "close": "last", + "volume": "sum", + "amount": "sum", + }).reset_index() + + result["trade_date"] = result["trade_time"].dt.date + result["freq"] = "2h" + result = result.iloc[offset:offset + limit] + return result + + # ---- Latest ---- + + def get_latest( + self, + freq: Freq, + ts_code: str | None = None, + ) -> pd.DataFrame: + """Get the latest K-line data for the most recent trade date.""" + if freq == Freq.h2: + return self._read_2h(ts_code, limit=100) + + glob = partition_glob(freq) + parquet_path = Path(settings.parquet_dir) / f"kline_{freq.storage_dir}" + if not parquet_path.exists(): + return pd.DataFrame(columns=STANDARD_COLS) + + conditions = [] + params = [] + if ts_code: + conditions.append(f"ts_code = ${len(params) + 1}") + params.append(ts_code) + + where_clause = "" + if conditions: + where_clause = "WHERE " + " AND ".join(conditions) + + sql = f""" + WITH latest AS ( + SELECT MAX(trade_date) AS max_date + FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true) + ) + SELECT k.* FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true) k + JOIN latest ON k.trade_date = latest.max_date + {where_clause} + ORDER BY k.ts_code + """ + try: + return self._db.conn.execute(sql, params).fetchdf() + except Exception as e: + logger.warning(f"get_latest failed for freq={freq.value}: {e}") + return pd.DataFrame(columns=STANDARD_COLS) + + # ---- Stats ---- + + def get_date_range(self, freq: Freq) -> tuple[date | None, date | None]: + """Get the min and max trade_date for a given frequency.""" + parquet_path = Path(settings.parquet_dir) / f"kline_{freq.storage_dir}" + if not parquet_path.exists(): + return None, None + glob = partition_glob(freq) + try: + row = self._db.conn.execute(f""" + SELECT MIN(trade_date), MAX(trade_date) + FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true) + """).fetchone() + return row[0], row[1] + except Exception: + return None, None + + def count_records(self, freq: Freq) -> int: + """Count total records for a given frequency.""" + parquet_path = Path(settings.parquet_dir) / f"kline_{freq.storage_dir}" + if not parquet_path.exists(): + return 0 + glob = partition_glob(freq) + try: + row = self._db.conn.execute(f""" + SELECT count(*) + FROM read_parquet('{glob}', hive_partitioning=true, union_by_name=true) + """).fetchone() + return row[0] + except Exception: + return 0 diff --git a/src/ashare_dp/storage/schema.py b/src/ashare_dp/storage/schema.py new file mode 100644 index 0000000..38ac3b1 --- /dev/null +++ b/src/ashare_dp/storage/schema.py @@ -0,0 +1,36 @@ +"""Database schema: DDL statements for DuckDB tables and views.""" + +from __future__ import annotations + +# DDL for persistent tables +DDL_STATEMENTS = [ + """ + CREATE TABLE IF NOT EXISTS stock_info ( + ts_code VARCHAR(9) PRIMARY KEY, + symbol VARCHAR(6) NOT NULL, + name VARCHAR(40) NOT NULL, + exchange VARCHAR(2) NOT NULL, + area VARCHAR(20), + industry VARCHAR(40), + list_date DATE, + delist_date DATE, + market VARCHAR(10), + updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + ) + """, + """ + CREATE INDEX IF NOT EXISTS idx_stock_symbol ON stock_info(symbol) + """, + """ + CREATE INDEX IF NOT EXISTS idx_stock_exchange ON stock_info(exchange) + """, + """ + CREATE TABLE IF NOT EXISTS trading_calendar ( + trade_date DATE PRIMARY KEY, + is_trading_day BOOLEAN NOT NULL DEFAULT TRUE, + week_day TINYINT NOT NULL, + year SMALLINT NOT NULL, + month TINYINT NOT NULL + ) + """, +]