quantdb-sdk 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- quantdb_sdk/__init__.py +33 -0
- quantdb_sdk/__main__.py +77 -0
- quantdb_sdk/_utils.py +68 -0
- quantdb_sdk/async_client.py +475 -0
- quantdb_sdk/client.py +610 -0
- quantdb_sdk/errors.py +43 -0
- quantdb_sdk/py.typed +0 -0
- quantdb_sdk-0.1.0.dist-info/METADATA +93 -0
- quantdb_sdk-0.1.0.dist-info/RECORD +11 -0
- quantdb_sdk-0.1.0.dist-info/WHEEL +5 -0
- quantdb_sdk-0.1.0.dist-info/top_level.txt +1 -0
quantdb_sdk/__init__.py
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
"""QuantDB 量化数据平台官方 Python SDK。"""
|
|
2
|
+
|
|
3
|
+
try:
|
|
4
|
+
from importlib.metadata import version
|
|
5
|
+
# 注意:此处用 PyPI 包名(连字符)查找版本号,不是模块名(下划线)
|
|
6
|
+
__version__ = version("quantdb-sdk")
|
|
7
|
+
except ImportError:
|
|
8
|
+
__version__ = "0.1.0"
|
|
9
|
+
|
|
10
|
+
from .async_client import AsyncQuantDBClient
|
|
11
|
+
from .client import DuckDBWarehouse, QuantDBClient
|
|
12
|
+
from .errors import (
|
|
13
|
+
AuthError,
|
|
14
|
+
InsufficientTrafficError,
|
|
15
|
+
NotFoundError,
|
|
16
|
+
QuantDBError,
|
|
17
|
+
RateLimitError,
|
|
18
|
+
ServerError,
|
|
19
|
+
ValidationError,
|
|
20
|
+
)
|
|
21
|
+
|
|
22
|
+
__all__ = [
|
|
23
|
+
"QuantDBClient",
|
|
24
|
+
"AsyncQuantDBClient",
|
|
25
|
+
"DuckDBWarehouse",
|
|
26
|
+
"QuantDBError",
|
|
27
|
+
"AuthError",
|
|
28
|
+
"InsufficientTrafficError",
|
|
29
|
+
"NotFoundError",
|
|
30
|
+
"RateLimitError",
|
|
31
|
+
"ServerError",
|
|
32
|
+
"ValidationError",
|
|
33
|
+
]
|
quantdb_sdk/__main__.py
ADDED
|
@@ -0,0 +1,77 @@
|
|
|
1
|
+
"""QuantDB SDK 命令行入口。
|
|
2
|
+
|
|
3
|
+
提供快速验证和诊断功能:
|
|
4
|
+
python -m quantdb --version
|
|
5
|
+
python -m quantdb --help
|
|
6
|
+
python -m quantdb check <api_key> # 验证 API Key 是否有效
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
import argparse
|
|
10
|
+
import sys
|
|
11
|
+
|
|
12
|
+
from . import __version__
|
|
13
|
+
from .client import QuantDBClient
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def main() -> int:
|
|
17
|
+
parser = argparse.ArgumentParser(
|
|
18
|
+
prog="quantdb",
|
|
19
|
+
description="QuantDB 量化数据平台官方 Python SDK",
|
|
20
|
+
)
|
|
21
|
+
parser.add_argument(
|
|
22
|
+
"--version",
|
|
23
|
+
action="version",
|
|
24
|
+
version=f"%(prog)s {__version__}",
|
|
25
|
+
)
|
|
26
|
+
parser.add_argument(
|
|
27
|
+
"--api-host",
|
|
28
|
+
default="https://quantdb.quantmind.cloud",
|
|
29
|
+
help="API 服务地址 (默认: https://quantdb.quantmind.cloud)",
|
|
30
|
+
)
|
|
31
|
+
|
|
32
|
+
subparsers = parser.add_subparsers(dest="command", help="可用命令")
|
|
33
|
+
|
|
34
|
+
# check 命令:验证 API Key
|
|
35
|
+
check_parser = subparsers.add_parser("check", help="验证 API Key 是否有效")
|
|
36
|
+
check_parser.add_argument("api_key", help="要验证的 API Key")
|
|
37
|
+
|
|
38
|
+
# usage 命令:查询用量
|
|
39
|
+
usage_parser = subparsers.add_parser("usage", help="查询账户用量和订阅状态")
|
|
40
|
+
usage_parser.add_argument("api_key", help="API Key")
|
|
41
|
+
|
|
42
|
+
args = parser.parse_args()
|
|
43
|
+
|
|
44
|
+
if not args.command:
|
|
45
|
+
parser.print_help()
|
|
46
|
+
return 0
|
|
47
|
+
|
|
48
|
+
try:
|
|
49
|
+
client = QuantDBClient(api_host=args.api_host, api_key=args.api_key)
|
|
50
|
+
|
|
51
|
+
if args.command == "check":
|
|
52
|
+
me = client.get_me()
|
|
53
|
+
print(f"✓ API Key 有效")
|
|
54
|
+
print(f" 用户: {me.get('username', 'N/A')}")
|
|
55
|
+
print(f" ID: {me.get('id', 'N/A')}")
|
|
56
|
+
return 0
|
|
57
|
+
|
|
58
|
+
if args.command == "usage":
|
|
59
|
+
usage = client.get_usage()
|
|
60
|
+
print(f"账户用量:")
|
|
61
|
+
print(f" 已用: {usage['used_gb']:.2f} GB")
|
|
62
|
+
print(f" 限额: {usage['limit_gb']:.1f} GB")
|
|
63
|
+
print(f" 剩余: {usage['remaining_gb']:.2f} GB")
|
|
64
|
+
print(f" 余额: ¥{usage['balance_yuan']:.2f}")
|
|
65
|
+
sub = usage.get("subscription", {})
|
|
66
|
+
print(f" 订阅: {sub.get('status', 'none')}")
|
|
67
|
+
return 0
|
|
68
|
+
|
|
69
|
+
except Exception as e:
|
|
70
|
+
print(f"✗ 错误: {e}", file=sys.stderr)
|
|
71
|
+
return 1
|
|
72
|
+
|
|
73
|
+
return 0
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
if __name__ == "__main__":
|
|
77
|
+
sys.exit(main())
|
quantdb_sdk/_utils.py
ADDED
|
@@ -0,0 +1,68 @@
|
|
|
1
|
+
"""QuantDB SDK 内部工具函数。
|
|
2
|
+
|
|
3
|
+
包含下载目录管理、文件名解析、字节换算等通用工具。
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
import os
|
|
7
|
+
import re
|
|
8
|
+
from urllib.parse import unquote
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def default_download_dir() -> str:
|
|
12
|
+
"""返回默认下载目录。
|
|
13
|
+
|
|
14
|
+
优先读取环境变量 ``QUANTDB_DOWNLOAD_DIR``,未设置时按平台回退:
|
|
15
|
+
|
|
16
|
+
- Windows: ``D:\\QuantDB\\downloads\\``
|
|
17
|
+
- Linux/macOS: ``~/QuantDB/downloads/``
|
|
18
|
+
|
|
19
|
+
Returns:
|
|
20
|
+
绝对路径字符串。
|
|
21
|
+
"""
|
|
22
|
+
env = os.getenv("QUANTDB_DOWNLOAD_DIR")
|
|
23
|
+
if env:
|
|
24
|
+
return env.strip()
|
|
25
|
+
if os.name == "nt":
|
|
26
|
+
return r"D:\QuantDB\downloads"
|
|
27
|
+
return os.path.join(os.path.expanduser("~"), "QuantDB", "downloads")
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def parse_filename_from_content_disposition(value: str, fallback: str) -> str:
|
|
31
|
+
"""从 HTTP Content-Disposition 响应头解析文件名。
|
|
32
|
+
|
|
33
|
+
支持 RFC 5987 (filename*=utf-8''name) 和常规 filename="name" 两种格式。
|
|
34
|
+
解析失败时返回 *fallback*。
|
|
35
|
+
|
|
36
|
+
Args:
|
|
37
|
+
value: Content-Disposition 头的值。
|
|
38
|
+
fallback: 解析失败时的默认文件名。
|
|
39
|
+
|
|
40
|
+
Returns:
|
|
41
|
+
解析后的文件名或 fallback。
|
|
42
|
+
"""
|
|
43
|
+
if not value:
|
|
44
|
+
return fallback
|
|
45
|
+
|
|
46
|
+
# RFC 5987 filename*=utf-8''name
|
|
47
|
+
m = re.search(r"filename\*\s*=\s*(?:[^']*)'" r"(?:[^']*)'([^;]+)", value, re.IGNORECASE)
|
|
48
|
+
if m:
|
|
49
|
+
return unquote(m.group(1))
|
|
50
|
+
|
|
51
|
+
# filename="name" 或 filename=name
|
|
52
|
+
m = re.search(r'filename\s*=\s*"?([^";]+)"?', value, re.IGNORECASE)
|
|
53
|
+
if m:
|
|
54
|
+
return unquote(m.group(1).strip('"'))
|
|
55
|
+
|
|
56
|
+
return fallback
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def bytes_to_gb(n: int) -> float:
|
|
60
|
+
"""将字节数转换为 GB(1024 进制)。
|
|
61
|
+
|
|
62
|
+
Args:
|
|
63
|
+
n: 字节数。
|
|
64
|
+
|
|
65
|
+
Returns:
|
|
66
|
+
GB 数值(浮点数)。
|
|
67
|
+
"""
|
|
68
|
+
return n / (1024 ** 3)
|
|
@@ -0,0 +1,475 @@
|
|
|
1
|
+
"""QuantDB 异步 Python SDK(基于 httpx)。"""
|
|
2
|
+
|
|
3
|
+
import io
|
|
4
|
+
import os
|
|
5
|
+
import re
|
|
6
|
+
from typing import Any, Dict, List, Optional
|
|
7
|
+
|
|
8
|
+
import httpx
|
|
9
|
+
import pandas as pd
|
|
10
|
+
|
|
11
|
+
from ._utils import bytes_to_gb, default_download_dir, parse_filename_from_content_disposition
|
|
12
|
+
from .errors import (
|
|
13
|
+
AuthError,
|
|
14
|
+
InsufficientTrafficError,
|
|
15
|
+
NotFoundError,
|
|
16
|
+
QuantDBError,
|
|
17
|
+
RateLimitError,
|
|
18
|
+
ServerError,
|
|
19
|
+
ValidationError,
|
|
20
|
+
)
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class AsyncQuantDBClient:
|
|
24
|
+
"""QuantDB 异步客户端。
|
|
25
|
+
|
|
26
|
+
基于 httpx,适用于 asyncio 量化框架。所有异步方法名以 a_ 开头。
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
def __init__(
|
|
30
|
+
self,
|
|
31
|
+
api_host: str = "https://quantdb.quantmind.cloud",
|
|
32
|
+
api_key: Optional[str] = None,
|
|
33
|
+
token: Optional[str] = None,
|
|
34
|
+
timeout: float = 60.0,
|
|
35
|
+
):
|
|
36
|
+
self.api_host = api_host.rstrip("/")
|
|
37
|
+
self.timeout = timeout
|
|
38
|
+
headers = {"User-Agent": "QuantDB-Python-SDK/0.1.0"}
|
|
39
|
+
if api_key:
|
|
40
|
+
headers["X-API-Key"] = api_key
|
|
41
|
+
elif token:
|
|
42
|
+
headers["Authorization"] = f"Bearer {token}"
|
|
43
|
+
else:
|
|
44
|
+
raise ValueError("必须提供 api_key 或 token")
|
|
45
|
+
|
|
46
|
+
self.client = httpx.AsyncClient(
|
|
47
|
+
headers=headers,
|
|
48
|
+
timeout=httpx.Timeout(timeout, connect=5.0),
|
|
49
|
+
)
|
|
50
|
+
# 进程内 parquet ETag 缓存:同对象(ETag 未变)不重复下载,避免重复计费。
|
|
51
|
+
self._cache: Dict[str, Dict[str, Any]] = {}
|
|
52
|
+
|
|
53
|
+
async def close(self) -> None:
|
|
54
|
+
"""关闭底层 httpx 客户端。"""
|
|
55
|
+
await self.client.aclose()
|
|
56
|
+
|
|
57
|
+
def clear_cache(self) -> None:
|
|
58
|
+
"""清空进程内 Parquet 缓存(强制下次重新下载最新数据)。"""
|
|
59
|
+
self._cache.clear()
|
|
60
|
+
|
|
61
|
+
async def __aenter__(self) -> "AsyncQuantDBClient":
|
|
62
|
+
return self
|
|
63
|
+
|
|
64
|
+
async def __aexit__(self, exc_type, exc, tb) -> None:
|
|
65
|
+
await self.close()
|
|
66
|
+
|
|
67
|
+
async def _request(
|
|
68
|
+
self,
|
|
69
|
+
method: str,
|
|
70
|
+
path: str,
|
|
71
|
+
params: Optional[dict] = None,
|
|
72
|
+
json: Optional[dict] = None,
|
|
73
|
+
) -> httpx.Response:
|
|
74
|
+
url = f"{self.api_host}{path}"
|
|
75
|
+
return await self.client.request(method=method, url=url, params=params, json=json)
|
|
76
|
+
|
|
77
|
+
async def _get(self, path: str, params: Optional[dict] = None) -> dict:
|
|
78
|
+
resp = await self._request("GET", path, params=params)
|
|
79
|
+
return self._check_response(resp)
|
|
80
|
+
|
|
81
|
+
async def _post(self, path: str, json: Optional[dict] = None) -> dict:
|
|
82
|
+
resp = await self._request("POST", path, json=json)
|
|
83
|
+
return self._check_response(resp)
|
|
84
|
+
|
|
85
|
+
def _check_response(self, response: httpx.Response) -> dict:
|
|
86
|
+
if response.status_code == 200:
|
|
87
|
+
try:
|
|
88
|
+
return response.json()
|
|
89
|
+
except Exception as exc:
|
|
90
|
+
raise QuantDBError(f"响应不是合法 JSON: {exc}") from exc
|
|
91
|
+
|
|
92
|
+
try:
|
|
93
|
+
payload = response.json()
|
|
94
|
+
detail = payload.get("detail", "")
|
|
95
|
+
except Exception:
|
|
96
|
+
detail = response.text[:200]
|
|
97
|
+
|
|
98
|
+
msg = detail or f"HTTP {response.status_code}"
|
|
99
|
+
|
|
100
|
+
if response.status_code == 401:
|
|
101
|
+
raise AuthError(f"认证失败:{msg}")
|
|
102
|
+
if response.status_code == 402:
|
|
103
|
+
raise InsufficientTrafficError(f"无有效订阅或本月流量已用完:{msg}")
|
|
104
|
+
if response.status_code == 403:
|
|
105
|
+
raise AuthError(f"无权访问:{msg}")
|
|
106
|
+
if response.status_code == 404:
|
|
107
|
+
raise NotFoundError(f"资源不存在:{msg}")
|
|
108
|
+
if response.status_code == 422:
|
|
109
|
+
raise ValidationError(f"参数错误:{msg}")
|
|
110
|
+
if response.status_code == 429:
|
|
111
|
+
raise RateLimitError(f"请求过于频繁:{msg}")
|
|
112
|
+
if response.status_code >= 500:
|
|
113
|
+
raise ServerError(f"服务端错误:{msg}")
|
|
114
|
+
raise QuantDBError(f"请求失败:{msg}")
|
|
115
|
+
|
|
116
|
+
# ========== 账户信息 ==========
|
|
117
|
+
|
|
118
|
+
async def a_get_me(self) -> Dict[str, Any]:
|
|
119
|
+
return await self._get("/api/v1/auth/me")
|
|
120
|
+
|
|
121
|
+
async def a_get_usage(self) -> Dict[str, Any]:
|
|
122
|
+
data = await self._get("/api/v1/auth/usage")
|
|
123
|
+
used = data.get("used_traffic", 0)
|
|
124
|
+
limit = data.get("traffic_limit", 0)
|
|
125
|
+
credit = data.get("credit_limit", 0)
|
|
126
|
+
subscription = data.get("subscription", {}) or {}
|
|
127
|
+
return {
|
|
128
|
+
"used_gb": bytes_to_gb(used),
|
|
129
|
+
"limit_gb": bytes_to_gb(limit),
|
|
130
|
+
"credit_gb": bytes_to_gb(credit),
|
|
131
|
+
"remaining_gb": max(0.0, bytes_to_gb(limit + credit - used)),
|
|
132
|
+
"balance_yuan": data.get("balance_yuan", 0.0),
|
|
133
|
+
"subscription": subscription,
|
|
134
|
+
"is_active": subscription.get("status") == "active" and used < limit + credit,
|
|
135
|
+
}
|
|
136
|
+
|
|
137
|
+
async def a_register(self, username: str, email: str, password: str) -> Dict[str, Any]:
|
|
138
|
+
return await self._post(
|
|
139
|
+
"/api/v1/auth/register",
|
|
140
|
+
json={"username": username, "email": email, "password": password},
|
|
141
|
+
)
|
|
142
|
+
|
|
143
|
+
async def a_forgot_password(self, email: str) -> Dict[str, Any]:
|
|
144
|
+
return await self._post("/api/v1/auth/forgot-password", json={"email": email})
|
|
145
|
+
|
|
146
|
+
async def a_reset_password(self, token: str, new_password: str) -> Dict[str, Any]:
|
|
147
|
+
return await self._post(
|
|
148
|
+
"/api/v1/auth/reset-password",
|
|
149
|
+
json={"token": token, "new_password": new_password},
|
|
150
|
+
)
|
|
151
|
+
|
|
152
|
+
# ========== API Key 管理 ==========
|
|
153
|
+
|
|
154
|
+
async def a_list_api_keys(self) -> List[Dict[str, Any]]:
|
|
155
|
+
data = await self._get("/api/v1/auth/api-keys")
|
|
156
|
+
return data.get("keys", data.get("data", []))
|
|
157
|
+
|
|
158
|
+
async def a_create_api_key(self, description: str = "") -> Dict[str, Any]:
|
|
159
|
+
"""创建新的 API Key(每个账号限制 1 个有效 Key)。"""
|
|
160
|
+
return await self._post("/api/v1/auth/api-keys", json={"description": description})
|
|
161
|
+
|
|
162
|
+
async def a_delete_api_key(self, key_token: str) -> Dict[str, Any]:
|
|
163
|
+
"""删除指定的 API Key。
|
|
164
|
+
|
|
165
|
+
Args:
|
|
166
|
+
key_token: 要删除的 API Key 令牌。
|
|
167
|
+
|
|
168
|
+
Returns:
|
|
169
|
+
删除结果信息。
|
|
170
|
+
"""
|
|
171
|
+
return await self._post(f"/api/v1/auth/api-keys/{key_token}/delete")
|
|
172
|
+
|
|
173
|
+
# ========== 订阅与支付 ==========
|
|
174
|
+
|
|
175
|
+
async def a_list_plans(self) -> List[Dict[str, Any]]:
|
|
176
|
+
data = await self._get("/api/v1/subscription/plans")
|
|
177
|
+
return data.get("plans", data.get("data", []))
|
|
178
|
+
|
|
179
|
+
async def a_create_order(self, plan_id: str) -> Dict[str, Any]:
|
|
180
|
+
return await self._post("/api/v1/subscription/orders", json={"plan_id": plan_id})
|
|
181
|
+
|
|
182
|
+
async def a_get_order(self, order_id: int) -> Dict[str, Any]:
|
|
183
|
+
return await self._get(f"/api/v1/subscription/orders/{order_id}")
|
|
184
|
+
|
|
185
|
+
async def a_get_version(self) -> Dict[str, Any]:
|
|
186
|
+
"""获取服务端版本信息。
|
|
187
|
+
|
|
188
|
+
Returns:
|
|
189
|
+
包含版本号、最新版本号、下载地址等信息的字典。
|
|
190
|
+
"""
|
|
191
|
+
return await self._get("/api/v1/version")
|
|
192
|
+
|
|
193
|
+
# ========== 数据查询:K线/Tick 走下载(消耗流量),列表/日历走网关 JSON(不计流量) ==========
|
|
194
|
+
|
|
195
|
+
async def a_query_kline(
|
|
196
|
+
self,
|
|
197
|
+
symbol: str,
|
|
198
|
+
adj_type: str = "unadjusted",
|
|
199
|
+
start_date: Optional[str] = None,
|
|
200
|
+
end_date: Optional[str] = None,
|
|
201
|
+
fields: str = "open,high,low,close,volume,amount",
|
|
202
|
+
limit: Optional[int] = None,
|
|
203
|
+
) -> pd.DataFrame:
|
|
204
|
+
"""查询 K 线数据(下载 COS parquet 切片后客户端解析,消耗下载流量,异步)。"""
|
|
205
|
+
sub_category = f"daily_{adj_type}"
|
|
206
|
+
df = await self.a_load_as_df("1", sub_category, symbol)
|
|
207
|
+
# 日期过滤:K线 parquet 的日期列为 'time',对外以 'trade_date' 暴露。
|
|
208
|
+
if "time" in df.columns:
|
|
209
|
+
dt = pd.to_datetime(df["time"], errors="coerce")
|
|
210
|
+
if dt.dt.tz is not None:
|
|
211
|
+
dt = dt.dt.tz_convert(None)
|
|
212
|
+
mask = pd.Series(True, index=df.index)
|
|
213
|
+
if start_date:
|
|
214
|
+
mask &= dt >= pd.to_datetime(start_date)
|
|
215
|
+
if end_date:
|
|
216
|
+
mask &= dt <= pd.to_datetime(end_date)
|
|
217
|
+
df = df[mask].copy()
|
|
218
|
+
df["trade_date"] = dt[mask].dt.strftime("%Y-%m-%d")
|
|
219
|
+
df = df.drop(columns=["time"])
|
|
220
|
+
field_list = [f.strip() for f in fields.split(",") if f.strip()]
|
|
221
|
+
cols = [c for c in field_list if c in df.columns]
|
|
222
|
+
keep = (["trade_date"] if "trade_date" in df.columns else []) + cols
|
|
223
|
+
keep = list(dict.fromkeys(keep))
|
|
224
|
+
df = df[keep]
|
|
225
|
+
if limit is not None:
|
|
226
|
+
df = df.tail(limit)
|
|
227
|
+
return df.reset_index(drop=True)
|
|
228
|
+
|
|
229
|
+
async def a_query_tick(
|
|
230
|
+
self,
|
|
231
|
+
symbol: str,
|
|
232
|
+
trade_date: str,
|
|
233
|
+
start_ts: Optional[str] = None,
|
|
234
|
+
end_ts: Optional[str] = None,
|
|
235
|
+
fields: str = "last_price,open,high,low,last_close,volume,amount",
|
|
236
|
+
limit: Optional[int] = None,
|
|
237
|
+
) -> pd.DataFrame:
|
|
238
|
+
"""查询 Tick 分笔数据(下载 COS parquet 切片后客户端解析,消耗下载流量,异步)。"""
|
|
239
|
+
df = await self.a_load_as_df("1", "tick_data", symbol, trade_date=trade_date)
|
|
240
|
+
ts_col = "ts" if "ts" in df.columns else ("time" if "time" in df.columns else None)
|
|
241
|
+
if ts_col and (start_ts or end_ts):
|
|
242
|
+
ts = pd.to_datetime(df[ts_col], errors="coerce")
|
|
243
|
+
if ts.dt.tz is not None:
|
|
244
|
+
ts = ts.dt.tz_convert(None)
|
|
245
|
+
mask = pd.Series(True, index=df.index)
|
|
246
|
+
if start_ts:
|
|
247
|
+
s = start_ts if (len(start_ts) > 8 or "-" in start_ts) else f"{trade_date} {start_ts}"
|
|
248
|
+
mask &= ts >= pd.to_datetime(s)
|
|
249
|
+
if end_ts:
|
|
250
|
+
e = end_ts if (len(end_ts) > 8 or "-" in end_ts) else f"{trade_date} {end_ts}"
|
|
251
|
+
mask &= ts <= pd.to_datetime(e)
|
|
252
|
+
df = df[mask].copy()
|
|
253
|
+
field_list = [f.strip() for f in fields.split(",") if f.strip()]
|
|
254
|
+
cols = [c for c in field_list if c in df.columns]
|
|
255
|
+
keep = ([ts_col] if ts_col else []) + cols
|
|
256
|
+
keep = list(dict.fromkeys(keep))
|
|
257
|
+
df = df[keep]
|
|
258
|
+
if limit is not None:
|
|
259
|
+
df = df.head(limit)
|
|
260
|
+
return df.reset_index(drop=True)
|
|
261
|
+
|
|
262
|
+
async def a_query_stock_list(
|
|
263
|
+
self, keyword: Optional[str] = None, limit: int = 200
|
|
264
|
+
) -> pd.DataFrame:
|
|
265
|
+
params = {"limit": limit}
|
|
266
|
+
if keyword:
|
|
267
|
+
params["keyword"] = keyword
|
|
268
|
+
data = await self._get("/api/v1/data/stock-list", params)
|
|
269
|
+
return pd.DataFrame(data.get("data", []))
|
|
270
|
+
|
|
271
|
+
async def a_query_calendar(
|
|
272
|
+
self,
|
|
273
|
+
start_date: Optional[str] = None,
|
|
274
|
+
end_date: Optional[str] = None,
|
|
275
|
+
) -> pd.DataFrame:
|
|
276
|
+
params: Dict[str, Any] = {}
|
|
277
|
+
if start_date:
|
|
278
|
+
params["start_date"] = start_date
|
|
279
|
+
if end_date:
|
|
280
|
+
params["end_date"] = end_date
|
|
281
|
+
data = await self._get("/api/v1/data/calendar", params)
|
|
282
|
+
return pd.DataFrame(data.get("data", []))
|
|
283
|
+
|
|
284
|
+
async def a_query_manifest(
|
|
285
|
+
self,
|
|
286
|
+
category_id: str,
|
|
287
|
+
sub_category: str,
|
|
288
|
+
trade_date: Optional[str] = None,
|
|
289
|
+
) -> List[Dict[str, Any]]:
|
|
290
|
+
params: Dict[str, Any] = {
|
|
291
|
+
"category_id": category_id,
|
|
292
|
+
"sub_category": sub_category,
|
|
293
|
+
}
|
|
294
|
+
if trade_date:
|
|
295
|
+
params["trade_date"] = trade_date
|
|
296
|
+
data = await self._get("/api/v1/data/download/manifest", params)
|
|
297
|
+
return data.get("files", [])
|
|
298
|
+
|
|
299
|
+
async def a_query_meta(
|
|
300
|
+
self,
|
|
301
|
+
dataset: Optional[str] = None,
|
|
302
|
+
source: Optional[str] = None,
|
|
303
|
+
) -> pd.DataFrame:
|
|
304
|
+
"""查询数据元数据(各数据集日期范围/行数/可用性,异步)。"""
|
|
305
|
+
params: Dict[str, Any] = {}
|
|
306
|
+
if dataset:
|
|
307
|
+
params["dataset"] = dataset
|
|
308
|
+
if source:
|
|
309
|
+
params["source"] = source
|
|
310
|
+
data = await self._get("/api/v1/data/meta", params)
|
|
311
|
+
return pd.DataFrame(data.get("data", []))
|
|
312
|
+
|
|
313
|
+
# ========== 数据预览与下载 ==========
|
|
314
|
+
|
|
315
|
+
async def a_preview_as_df(
|
|
316
|
+
self,
|
|
317
|
+
category_id: str,
|
|
318
|
+
sub_category: str,
|
|
319
|
+
symbol: Optional[str] = None,
|
|
320
|
+
limit: int = 30,
|
|
321
|
+
) -> pd.DataFrame:
|
|
322
|
+
params: Dict[str, Any] = {
|
|
323
|
+
"category_id": category_id,
|
|
324
|
+
"sub_category": sub_category,
|
|
325
|
+
"limit": limit,
|
|
326
|
+
}
|
|
327
|
+
if symbol:
|
|
328
|
+
params["symbol"] = symbol
|
|
329
|
+
data = await self._get("/api/v1/data/preview", params)
|
|
330
|
+
return pd.DataFrame(data.get("data", []), columns=data.get("columns"))
|
|
331
|
+
|
|
332
|
+
async def a_download_file(
|
|
333
|
+
self,
|
|
334
|
+
category_id: str,
|
|
335
|
+
sub_category: str,
|
|
336
|
+
symbol: Optional[str] = None,
|
|
337
|
+
save_dir: Optional[str] = None,
|
|
338
|
+
trade_date: Optional[str] = None,
|
|
339
|
+
) -> str:
|
|
340
|
+
if save_dir is None:
|
|
341
|
+
save_dir = default_download_dir()
|
|
342
|
+
os.makedirs(save_dir, exist_ok=True)
|
|
343
|
+
|
|
344
|
+
params: Dict[str, Any] = {
|
|
345
|
+
"category_id": category_id,
|
|
346
|
+
"sub_category": sub_category,
|
|
347
|
+
}
|
|
348
|
+
if symbol:
|
|
349
|
+
params["symbol"] = symbol
|
|
350
|
+
if trade_date:
|
|
351
|
+
params["trade_date"] = trade_date
|
|
352
|
+
|
|
353
|
+
async with self.client.stream(
|
|
354
|
+
"GET", f"{self.api_host}/api/v1/data/download", params=params
|
|
355
|
+
) as resp:
|
|
356
|
+
if resp.status_code != 200:
|
|
357
|
+
body = await resp.aread()
|
|
358
|
+
# 构造一个完整的 httpx.Response 用于错误解析,然后立即抛出
|
|
359
|
+
err_resp = httpx.Response(resp.status_code, content=body)
|
|
360
|
+
self._check_response(err_resp)
|
|
361
|
+
# 如果 _check_response 没有抛出(理论上不应该),安全退出
|
|
362
|
+
raise QuantDBError(f"下载失败,HTTP {resp.status_code}")
|
|
363
|
+
|
|
364
|
+
cd = ""
|
|
365
|
+
if "Content-Disposition" in resp.headers:
|
|
366
|
+
cd = resp.headers["Content-Disposition"]
|
|
367
|
+
fallback = f"{sub_category}.parquet"
|
|
368
|
+
if symbol:
|
|
369
|
+
fallback = f"{sub_category}_{symbol}.parquet"
|
|
370
|
+
filename = parse_filename_from_content_disposition(cd, fallback)
|
|
371
|
+
|
|
372
|
+
save_path = os.path.join(save_dir, filename)
|
|
373
|
+
with open(save_path, "wb") as f:
|
|
374
|
+
async for chunk in resp.aiter_bytes(chunk_size=8192):
|
|
375
|
+
if chunk:
|
|
376
|
+
f.write(chunk)
|
|
377
|
+
return os.path.abspath(save_path)
|
|
378
|
+
|
|
379
|
+
async def a_load_as_df(
|
|
380
|
+
self,
|
|
381
|
+
category_id: str,
|
|
382
|
+
sub_category: str,
|
|
383
|
+
symbol: Optional[str] = None,
|
|
384
|
+
trade_date: Optional[str] = None,
|
|
385
|
+
) -> pd.DataFrame:
|
|
386
|
+
"""远端 Parquet 切片加载到内存 DataFrame(消耗下载流量,异步)。
|
|
387
|
+
|
|
388
|
+
带进程内 ETag 缓存:同对象(ETag 未变)不重复下载,避免重复计费。
|
|
389
|
+
"""
|
|
390
|
+
cache_key = f"{category_id}/{sub_category}/{symbol}/{trade_date}"
|
|
391
|
+
params: Dict[str, Any] = {
|
|
392
|
+
"category_id": category_id,
|
|
393
|
+
"sub_category": sub_category,
|
|
394
|
+
}
|
|
395
|
+
if symbol:
|
|
396
|
+
params["symbol"] = symbol
|
|
397
|
+
if trade_date:
|
|
398
|
+
params["trade_date"] = trade_date
|
|
399
|
+
cached = self._cache.get(cache_key)
|
|
400
|
+
req_headers = {}
|
|
401
|
+
if cached and cached.get("etag"):
|
|
402
|
+
req_headers["If-None-Match"] = cached["etag"]
|
|
403
|
+
resp = await self.client.get(
|
|
404
|
+
f"{self.api_host}/api/v1/data/download", params=params, headers=req_headers
|
|
405
|
+
)
|
|
406
|
+
if resp.status_code == 304 and cached:
|
|
407
|
+
return cached["df"]
|
|
408
|
+
if resp.status_code != 200:
|
|
409
|
+
self._check_response(resp)
|
|
410
|
+
etag = resp.headers.get("ETag") or ""
|
|
411
|
+
if cached and etag and cached.get("etag") == etag:
|
|
412
|
+
return cached["df"]
|
|
413
|
+
df = pd.read_parquet(io.BytesIO(resp.content))
|
|
414
|
+
if etag:
|
|
415
|
+
self._cache[cache_key] = {"etag": etag, "df": df}
|
|
416
|
+
return df
|
|
417
|
+
|
|
418
|
+
async def a_query_local(
|
|
419
|
+
self,
|
|
420
|
+
sql: str,
|
|
421
|
+
category_id: str,
|
|
422
|
+
sub_category: str,
|
|
423
|
+
symbol: Optional[str] = None,
|
|
424
|
+
save_dir: Optional[str] = None,
|
|
425
|
+
) -> pd.DataFrame:
|
|
426
|
+
"""用 DuckDB 查询本地已下载的 Parquet 文件(零流量,异步)。
|
|
427
|
+
|
|
428
|
+
若文件不存在会先自动下载。
|
|
429
|
+
|
|
430
|
+
sql 支持两种写法:
|
|
431
|
+
1. 完整 SELECT 语句(必须包含 FROM 子句,表名会被替换为 parquet 路径)
|
|
432
|
+
2. 纯 WHERE 条件字符串(不含 FROM 时自动拼装为 SELECT * FROM <path> WHERE <条件>)
|
|
433
|
+
|
|
434
|
+
安全限制:
|
|
435
|
+
- 不允许分号(;),防止多语句注入
|
|
436
|
+
- 不允许注释(-- 或 /* */),防止注释注入
|
|
437
|
+
- 不允许 UNION / DROP / DELETE / INSERT / UPDATE / ALTER / CREATE / EXEC / EXECUTE
|
|
438
|
+
"""
|
|
439
|
+
file_path = await self.a_download_file(
|
|
440
|
+
category_id=category_id,
|
|
441
|
+
sub_category=sub_category,
|
|
442
|
+
symbol=symbol,
|
|
443
|
+
save_dir=save_dir,
|
|
444
|
+
)
|
|
445
|
+
clean_path = file_path.replace("\\", "/")
|
|
446
|
+
|
|
447
|
+
try:
|
|
448
|
+
import duckdb # type: ignore
|
|
449
|
+
except ImportError as exc:
|
|
450
|
+
raise RuntimeError("本地 SQL 查询需要 duckdb:pip install duckdb") from exc
|
|
451
|
+
|
|
452
|
+
# 安全检查:拒绝危险字符和关键字
|
|
453
|
+
dangerous_keywords = [
|
|
454
|
+
";", "--", "/*", "*/", "union", "drop", "delete",
|
|
455
|
+
"insert", "update", "alter", "create", "exec", "execute",
|
|
456
|
+
]
|
|
457
|
+
sql_upper = sql.upper()
|
|
458
|
+
for kw in dangerous_keywords:
|
|
459
|
+
if kw in sql_upper:
|
|
460
|
+
raise ValidationError(
|
|
461
|
+
f"SQL 包含危险关键字 '{kw}',已被拒绝。"
|
|
462
|
+
"a_query_local 仅支持 SELECT 查询,不支持数据修改或联合查询。"
|
|
463
|
+
)
|
|
464
|
+
|
|
465
|
+
if "FROM" in sql_upper:
|
|
466
|
+
sql = re.sub(
|
|
467
|
+
r"FROM\s+[`\"\']?\w+[`\"\']?",
|
|
468
|
+
f"FROM '{clean_path}'",
|
|
469
|
+
sql,
|
|
470
|
+
count=1,
|
|
471
|
+
flags=re.IGNORECASE,
|
|
472
|
+
)
|
|
473
|
+
else:
|
|
474
|
+
sql = f"SELECT * FROM '{clean_path}' WHERE {sql}"
|
|
475
|
+
return duckdb.query(sql).df()
|