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/client.py
ADDED
|
@@ -0,0 +1,610 @@
|
|
|
1
|
+
"""QuantDB 同步 Python SDK。"""
|
|
2
|
+
|
|
3
|
+
import io
|
|
4
|
+
import os
|
|
5
|
+
import re
|
|
6
|
+
from typing import Any, Dict, List, Optional
|
|
7
|
+
|
|
8
|
+
import pandas as pd
|
|
9
|
+
import requests
|
|
10
|
+
from requests.adapters import HTTPAdapter
|
|
11
|
+
from urllib3.util.retry import Retry
|
|
12
|
+
|
|
13
|
+
from ._utils import bytes_to_gb, default_download_dir, parse_filename_from_content_disposition
|
|
14
|
+
from .errors import (
|
|
15
|
+
AuthError,
|
|
16
|
+
InsufficientTrafficError,
|
|
17
|
+
NotFoundError,
|
|
18
|
+
QuantDBError,
|
|
19
|
+
RateLimitError,
|
|
20
|
+
ServerError,
|
|
21
|
+
ValidationError,
|
|
22
|
+
)
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class QuantDBClient:
|
|
26
|
+
"""QuantDB 同步客户端。
|
|
27
|
+
|
|
28
|
+
支持 API Key 或用户名/密码两种鉴权方式。
|
|
29
|
+
数据查询:K线/Tick 走下载(消耗流量),列表/日历/元数据走网关 JSON(不计流量)。
|
|
30
|
+
"""
|
|
31
|
+
|
|
32
|
+
def __init__(
|
|
33
|
+
self,
|
|
34
|
+
api_host: str = "https://quantdb.quantmind.cloud",
|
|
35
|
+
api_key: Optional[str] = None,
|
|
36
|
+
username: Optional[str] = None,
|
|
37
|
+
password: Optional[str] = None,
|
|
38
|
+
timeout: tuple = (5, 60),
|
|
39
|
+
max_retries: int = 2,
|
|
40
|
+
):
|
|
41
|
+
self.api_host = api_host.rstrip("/")
|
|
42
|
+
self.timeout = timeout
|
|
43
|
+
self.token: Optional[str] = None
|
|
44
|
+
self.headers: Dict[str, str] = {}
|
|
45
|
+
|
|
46
|
+
if api_key:
|
|
47
|
+
self.headers = {"X-API-Key": api_key}
|
|
48
|
+
elif username and password:
|
|
49
|
+
self._login(username, password)
|
|
50
|
+
else:
|
|
51
|
+
raise ValueError("必须提供 api_key 或 username+password")
|
|
52
|
+
|
|
53
|
+
self.session = requests.Session()
|
|
54
|
+
self.session.headers.update({"User-Agent": "QuantDB-Python-SDK/0.1.0"})
|
|
55
|
+
if self.headers:
|
|
56
|
+
self.session.headers.update(self.headers)
|
|
57
|
+
|
|
58
|
+
# 进程内 parquet 缓存:cache_key -> {"etag": str, "df": DataFrame}
|
|
59
|
+
# 用 ETag 判断对象是否变化:同对象(同 ETag)不重复下载,避免重复计费。
|
|
60
|
+
# 典型场景:量化脚本内对同一 symbol 多次查询不同字段/日期范围,底层 parquet
|
|
61
|
+
# 只需下载一次。进程退出即失效(无跨进程副作用)。
|
|
62
|
+
self._cache: Dict[str, Dict[str, Any]] = {}
|
|
63
|
+
|
|
64
|
+
retries = Retry(
|
|
65
|
+
total=max_retries,
|
|
66
|
+
backoff_factor=0.5,
|
|
67
|
+
status_forcelist=[429, 502, 503, 504],
|
|
68
|
+
allowed_methods=["GET", "POST"],
|
|
69
|
+
)
|
|
70
|
+
self.session.mount("http://", HTTPAdapter(max_retries=retries))
|
|
71
|
+
self.session.mount("https://", HTTPAdapter(max_retries=retries))
|
|
72
|
+
|
|
73
|
+
def _login(self, username: str, password: str) -> None:
|
|
74
|
+
"""用户名密码登录,获取 JWT Token。"""
|
|
75
|
+
resp = self.session.post(
|
|
76
|
+
f"{self.api_host}/api/v1/auth/login",
|
|
77
|
+
json={"username": username, "password": password},
|
|
78
|
+
timeout=self.timeout,
|
|
79
|
+
)
|
|
80
|
+
data = self._check_response(resp)
|
|
81
|
+
self.token = data["access_token"]
|
|
82
|
+
self.session.headers.update({"Authorization": f"Bearer {self.token}"})
|
|
83
|
+
|
|
84
|
+
def _request(
|
|
85
|
+
self,
|
|
86
|
+
method: str,
|
|
87
|
+
path: str,
|
|
88
|
+
params: Optional[dict] = None,
|
|
89
|
+
json: Optional[dict] = None,
|
|
90
|
+
stream: bool = False,
|
|
91
|
+
headers: Optional[dict] = None,
|
|
92
|
+
) -> requests.Response:
|
|
93
|
+
url = f"{self.api_host}{path}"
|
|
94
|
+
resp = self.session.request(
|
|
95
|
+
method=method,
|
|
96
|
+
url=url,
|
|
97
|
+
params=params,
|
|
98
|
+
json=json,
|
|
99
|
+
stream=stream,
|
|
100
|
+
timeout=self.timeout,
|
|
101
|
+
headers=headers,
|
|
102
|
+
)
|
|
103
|
+
return resp
|
|
104
|
+
|
|
105
|
+
def _get(self, path: str, params: Optional[dict] = None) -> dict:
|
|
106
|
+
resp = self._request("GET", path, params=params)
|
|
107
|
+
return self._check_response(resp)
|
|
108
|
+
|
|
109
|
+
def _post(self, path: str, json: Optional[dict] = None) -> dict:
|
|
110
|
+
resp = self._request("POST", path, json=json)
|
|
111
|
+
return self._check_response(resp)
|
|
112
|
+
|
|
113
|
+
def _check_response(self, response: requests.Response) -> dict:
|
|
114
|
+
if response.status_code == 200:
|
|
115
|
+
try:
|
|
116
|
+
return response.json()
|
|
117
|
+
except Exception as exc:
|
|
118
|
+
raise QuantDBError(f"响应不是合法 JSON: {exc}") from exc
|
|
119
|
+
|
|
120
|
+
try:
|
|
121
|
+
payload = response.json()
|
|
122
|
+
detail = payload.get("detail", "")
|
|
123
|
+
except Exception:
|
|
124
|
+
detail = response.text[:200]
|
|
125
|
+
|
|
126
|
+
msg = detail or f"HTTP {response.status_code}"
|
|
127
|
+
|
|
128
|
+
if response.status_code == 401:
|
|
129
|
+
raise AuthError(f"认证失败:{msg}")
|
|
130
|
+
if response.status_code == 402:
|
|
131
|
+
raise InsufficientTrafficError(f"无有效订阅或本月流量已用完:{msg}")
|
|
132
|
+
if response.status_code == 403:
|
|
133
|
+
raise AuthError(f"无权访问:{msg}")
|
|
134
|
+
if response.status_code == 404:
|
|
135
|
+
raise NotFoundError(f"资源不存在:{msg}")
|
|
136
|
+
if response.status_code == 422:
|
|
137
|
+
raise ValidationError(f"参数错误:{msg}")
|
|
138
|
+
if response.status_code == 429:
|
|
139
|
+
raise RateLimitError(f"请求过于频繁:{msg}")
|
|
140
|
+
if response.status_code >= 500:
|
|
141
|
+
raise ServerError(f"服务端错误:{msg}")
|
|
142
|
+
raise QuantDBError(f"请求失败:{msg}")
|
|
143
|
+
|
|
144
|
+
def clear_cache(self) -> None:
|
|
145
|
+
"""清空进程内 Parquet 缓存。
|
|
146
|
+
|
|
147
|
+
load_as_df 会按 ETag 缓存已下载的 parquet,避免重复查询同 symbol 被重复计费。
|
|
148
|
+
当确认需要强制重新下载最新数据时调用此方法。
|
|
149
|
+
"""
|
|
150
|
+
self._cache.clear()
|
|
151
|
+
|
|
152
|
+
# ========== 账户信息 ==========
|
|
153
|
+
|
|
154
|
+
def get_me(self) -> Dict[str, Any]:
|
|
155
|
+
"""获取当前登录用户信息。"""
|
|
156
|
+
return self._get("/api/v1/auth/me")
|
|
157
|
+
|
|
158
|
+
def get_usage(self) -> Dict[str, Any]:
|
|
159
|
+
"""获取当前账户流量与订阅使用情况。"""
|
|
160
|
+
data = self._get("/api/v1/auth/usage")
|
|
161
|
+
used = data.get("used_traffic", 0)
|
|
162
|
+
limit = data.get("traffic_limit", 0)
|
|
163
|
+
credit = data.get("credit_limit", 0)
|
|
164
|
+
subscription = data.get("subscription", {}) or {}
|
|
165
|
+
return {
|
|
166
|
+
"used_gb": bytes_to_gb(used),
|
|
167
|
+
"limit_gb": bytes_to_gb(limit),
|
|
168
|
+
"credit_gb": bytes_to_gb(credit),
|
|
169
|
+
"remaining_gb": max(0.0, bytes_to_gb(limit + credit - used)),
|
|
170
|
+
"balance_yuan": data.get("balance_yuan", 0.0),
|
|
171
|
+
"subscription": subscription,
|
|
172
|
+
"is_active": subscription.get("status") == "active" and used < limit + credit,
|
|
173
|
+
}
|
|
174
|
+
|
|
175
|
+
def register(self, username: str, email: str, password: str) -> Dict[str, Any]:
|
|
176
|
+
"""用户注册。"""
|
|
177
|
+
return self._post(
|
|
178
|
+
"/api/v1/auth/register",
|
|
179
|
+
json={"username": username, "email": email, "password": password},
|
|
180
|
+
)
|
|
181
|
+
|
|
182
|
+
def forgot_password(self, email: str) -> Dict[str, Any]:
|
|
183
|
+
"""发送密码重置邮件。"""
|
|
184
|
+
return self._post("/api/v1/auth/forgot-password", json={"email": email})
|
|
185
|
+
|
|
186
|
+
def reset_password(self, token: str, new_password: str) -> Dict[str, Any]:
|
|
187
|
+
"""使用邮件中的重置令牌设置新密码。"""
|
|
188
|
+
return self._post(
|
|
189
|
+
"/api/v1/auth/reset-password",
|
|
190
|
+
json={"token": token, "new_password": new_password},
|
|
191
|
+
)
|
|
192
|
+
|
|
193
|
+
# ========== API Key 管理 ==========
|
|
194
|
+
|
|
195
|
+
def list_api_keys(self) -> List[Dict[str, Any]]:
|
|
196
|
+
"""列出当前用户的 API Key。"""
|
|
197
|
+
data = self._get("/api/v1/auth/api-keys")
|
|
198
|
+
return data.get("keys", data.get("data", []))
|
|
199
|
+
|
|
200
|
+
def create_api_key(self, description: str = "") -> Dict[str, Any]:
|
|
201
|
+
"""创建新的 API Key(每个账号限制 1 个有效 Key)。"""
|
|
202
|
+
return self._post("/api/v1/auth/api-keys", json={"description": description})
|
|
203
|
+
|
|
204
|
+
def delete_api_key(self, key_token: str) -> Dict[str, Any]:
|
|
205
|
+
"""删除指定的 API Key。
|
|
206
|
+
|
|
207
|
+
Args:
|
|
208
|
+
key_token: 要删除的 API Key 令牌。
|
|
209
|
+
|
|
210
|
+
Returns:
|
|
211
|
+
删除结果信息。
|
|
212
|
+
"""
|
|
213
|
+
return self._post(f"/api/v1/auth/api-keys/{key_token}/delete")
|
|
214
|
+
|
|
215
|
+
# ========== 订阅与支付 ==========
|
|
216
|
+
|
|
217
|
+
def list_plans(self) -> List[Dict[str, Any]]:
|
|
218
|
+
"""获取可购买套餐列表。"""
|
|
219
|
+
data = self._get("/api/v1/subscription/plans")
|
|
220
|
+
return data.get("plans", data.get("data", []))
|
|
221
|
+
|
|
222
|
+
def create_order(self, plan_id: str) -> Dict[str, Any]:
|
|
223
|
+
"""创建订单,返回支付跳转 URL。"""
|
|
224
|
+
return self._post("/api/v1/subscription/orders", json={"plan_id": plan_id})
|
|
225
|
+
|
|
226
|
+
def get_order(self, order_id: int) -> Dict[str, Any]:
|
|
227
|
+
"""查询订单状态。"""
|
|
228
|
+
return self._get(f"/api/v1/subscription/orders/{order_id}")
|
|
229
|
+
|
|
230
|
+
def get_version(self) -> Dict[str, Any]:
|
|
231
|
+
"""获取服务端版本信息。
|
|
232
|
+
|
|
233
|
+
Returns:
|
|
234
|
+
包含版本号、最新版本号、下载地址等信息的字典。
|
|
235
|
+
"""
|
|
236
|
+
return self._get("/api/v1/version")
|
|
237
|
+
|
|
238
|
+
# ========== 数据查询:K线/Tick 走下载(消耗流量),列表/日历走网关 JSON(不计流量) ==========
|
|
239
|
+
|
|
240
|
+
def query_kline(
|
|
241
|
+
self,
|
|
242
|
+
symbol: str,
|
|
243
|
+
adj_type: str = "unadjusted",
|
|
244
|
+
start_date: Optional[str] = None,
|
|
245
|
+
end_date: Optional[str] = None,
|
|
246
|
+
fields: str = "open,high,low,close,volume,amount",
|
|
247
|
+
limit: Optional[int] = None,
|
|
248
|
+
) -> pd.DataFrame:
|
|
249
|
+
"""查询 K 线数据(下载 COS parquet 切片后客户端解析,消耗下载流量)。
|
|
250
|
+
|
|
251
|
+
下载整个 symbol 的日线 parquet,按 start_date/end_date 过滤日期、按 fields 选列。
|
|
252
|
+
limit 非 None 时只返回尾部 limit 行。
|
|
253
|
+
"""
|
|
254
|
+
sub_category = f"daily_{adj_type}"
|
|
255
|
+
df = self.load_as_df("1", sub_category, symbol)
|
|
256
|
+
# 日期过滤:K线 parquet 的日期列为 'time'(采集端 01_kline.py 统一命名)。
|
|
257
|
+
# 对外仍以 'trade_date' 暴露(YYYY-MM-DD),便于用户按日期理解。
|
|
258
|
+
if "time" in df.columns:
|
|
259
|
+
dt = pd.to_datetime(df["time"], errors="coerce")
|
|
260
|
+
if dt.dt.tz is not None:
|
|
261
|
+
dt = dt.dt.tz_convert(None)
|
|
262
|
+
mask = pd.Series(True, index=df.index)
|
|
263
|
+
if start_date:
|
|
264
|
+
mask &= dt >= pd.to_datetime(start_date)
|
|
265
|
+
if end_date:
|
|
266
|
+
mask &= dt <= pd.to_datetime(end_date)
|
|
267
|
+
df = df[mask].copy()
|
|
268
|
+
df["trade_date"] = dt[mask].dt.strftime("%Y-%m-%d")
|
|
269
|
+
df = df.drop(columns=["time"])
|
|
270
|
+
# 字段过滤:保留 trade_date + 选中字段
|
|
271
|
+
field_list = [f.strip() for f in fields.split(",") if f.strip()]
|
|
272
|
+
cols = [c for c in field_list if c in df.columns]
|
|
273
|
+
keep = (["trade_date"] if "trade_date" in df.columns else []) + cols
|
|
274
|
+
keep = list(dict.fromkeys(keep))
|
|
275
|
+
df = df[keep]
|
|
276
|
+
if limit is not None:
|
|
277
|
+
df = df.tail(limit)
|
|
278
|
+
return df.reset_index(drop=True)
|
|
279
|
+
|
|
280
|
+
def query_tick(
|
|
281
|
+
self,
|
|
282
|
+
symbol: str,
|
|
283
|
+
trade_date: str,
|
|
284
|
+
start_ts: Optional[str] = None,
|
|
285
|
+
end_ts: Optional[str] = None,
|
|
286
|
+
fields: str = "last_price,open,high,low,last_close,volume,amount",
|
|
287
|
+
limit: Optional[int] = None,
|
|
288
|
+
) -> pd.DataFrame:
|
|
289
|
+
"""查询 Tick 分笔数据(下载 COS parquet 切片后客户端解析,消耗下载流量)。
|
|
290
|
+
|
|
291
|
+
下载 trade_date 当日该 symbol 的 tick parquet,按 start_ts/end_ts 过滤时间、按 fields 选列。
|
|
292
|
+
start_ts/end_ts 可传完整时间戳或 "HH:MM:SS"(自动补 trade_date 日期)。
|
|
293
|
+
"""
|
|
294
|
+
df = self.load_as_df("1", "tick_data", symbol, trade_date=trade_date)
|
|
295
|
+
# 时间过滤
|
|
296
|
+
ts_col = "ts" if "ts" in df.columns else ("time" if "time" in df.columns else None)
|
|
297
|
+
if ts_col and (start_ts or end_ts):
|
|
298
|
+
ts = pd.to_datetime(df[ts_col], errors="coerce")
|
|
299
|
+
if ts.dt.tz is not None:
|
|
300
|
+
ts = ts.dt.tz_convert(None)
|
|
301
|
+
mask = pd.Series(True, index=df.index)
|
|
302
|
+
if start_ts:
|
|
303
|
+
s = start_ts if (len(start_ts) > 8 or "-" in start_ts) else f"{trade_date} {start_ts}"
|
|
304
|
+
mask &= ts >= pd.to_datetime(s)
|
|
305
|
+
if end_ts:
|
|
306
|
+
e = end_ts if (len(end_ts) > 8 or "-" in end_ts) else f"{trade_date} {end_ts}"
|
|
307
|
+
mask &= ts <= pd.to_datetime(e)
|
|
308
|
+
df = df[mask].copy()
|
|
309
|
+
# 字段过滤
|
|
310
|
+
field_list = [f.strip() for f in fields.split(",") if f.strip()]
|
|
311
|
+
cols = [c for c in field_list if c in df.columns]
|
|
312
|
+
keep = ([ts_col] if ts_col else []) + cols
|
|
313
|
+
keep = list(dict.fromkeys(keep))
|
|
314
|
+
df = df[keep]
|
|
315
|
+
if limit is not None:
|
|
316
|
+
df = df.head(limit)
|
|
317
|
+
return df.reset_index(drop=True)
|
|
318
|
+
|
|
319
|
+
def query_stock_list(
|
|
320
|
+
self, keyword: Optional[str] = None, limit: int = 200
|
|
321
|
+
) -> pd.DataFrame:
|
|
322
|
+
"""查询 A 股基础列表。"""
|
|
323
|
+
params = {"limit": limit}
|
|
324
|
+
if keyword:
|
|
325
|
+
params["keyword"] = keyword
|
|
326
|
+
data = self._get("/api/v1/data/stock-list", params)
|
|
327
|
+
return pd.DataFrame(data.get("data", []))
|
|
328
|
+
|
|
329
|
+
def query_calendar(
|
|
330
|
+
self,
|
|
331
|
+
start_date: Optional[str] = None,
|
|
332
|
+
end_date: Optional[str] = None,
|
|
333
|
+
) -> pd.DataFrame:
|
|
334
|
+
"""查询 A 股交易日历。"""
|
|
335
|
+
params: Dict[str, Any] = {}
|
|
336
|
+
if start_date:
|
|
337
|
+
params["start_date"] = start_date
|
|
338
|
+
if end_date:
|
|
339
|
+
params["end_date"] = end_date
|
|
340
|
+
data = self._get("/api/v1/data/calendar", params)
|
|
341
|
+
return pd.DataFrame(data.get("data", []))
|
|
342
|
+
|
|
343
|
+
def query_manifest(
|
|
344
|
+
self,
|
|
345
|
+
category_id: str,
|
|
346
|
+
sub_category: str,
|
|
347
|
+
trade_date: Optional[str] = None,
|
|
348
|
+
) -> List[Dict[str, Any]]:
|
|
349
|
+
"""查询 COS 可下载文件清单。"""
|
|
350
|
+
params: Dict[str, Any] = {
|
|
351
|
+
"category_id": category_id,
|
|
352
|
+
"sub_category": sub_category,
|
|
353
|
+
}
|
|
354
|
+
if trade_date:
|
|
355
|
+
params["trade_date"] = trade_date
|
|
356
|
+
data = self._get("/api/v1/data/download/manifest", params)
|
|
357
|
+
return data.get("files", [])
|
|
358
|
+
|
|
359
|
+
def query_meta(
|
|
360
|
+
self,
|
|
361
|
+
dataset: Optional[str] = None,
|
|
362
|
+
source: Optional[str] = None,
|
|
363
|
+
) -> pd.DataFrame:
|
|
364
|
+
"""查询数据元数据(各数据集日期范围/行数/可用性)。
|
|
365
|
+
|
|
366
|
+
数据来自 PostgreSQL data_meta 表(由每日流水线从 COS 重建)。
|
|
367
|
+
可按 dataset(kline/financial/basic/calendar/sector/tick)与
|
|
368
|
+
source(cos)过滤。
|
|
369
|
+
"""
|
|
370
|
+
params: Dict[str, Any] = {}
|
|
371
|
+
if dataset:
|
|
372
|
+
params["dataset"] = dataset
|
|
373
|
+
if source:
|
|
374
|
+
params["source"] = source
|
|
375
|
+
data = self._get("/api/v1/data/meta", params)
|
|
376
|
+
return pd.DataFrame(data.get("data", []))
|
|
377
|
+
|
|
378
|
+
# ========== 数据预览与下载(消耗下载流量) ==========
|
|
379
|
+
|
|
380
|
+
def preview_as_df(
|
|
381
|
+
self,
|
|
382
|
+
category_id: str,
|
|
383
|
+
sub_category: str,
|
|
384
|
+
symbol: Optional[str] = None,
|
|
385
|
+
limit: int = 30,
|
|
386
|
+
) -> pd.DataFrame:
|
|
387
|
+
"""预览 Parquet 尾部 N 条数据(不消耗流量)。"""
|
|
388
|
+
params: Dict[str, Any] = {
|
|
389
|
+
"category_id": category_id,
|
|
390
|
+
"sub_category": sub_category,
|
|
391
|
+
"limit": limit,
|
|
392
|
+
}
|
|
393
|
+
if symbol:
|
|
394
|
+
params["symbol"] = symbol
|
|
395
|
+
data = self._get("/api/v1/data/preview", params)
|
|
396
|
+
return pd.DataFrame(data.get("data", []), columns=data.get("columns"))
|
|
397
|
+
|
|
398
|
+
def download_file(
|
|
399
|
+
self,
|
|
400
|
+
category_id: str,
|
|
401
|
+
sub_category: str,
|
|
402
|
+
symbol: Optional[str] = None,
|
|
403
|
+
save_dir: Optional[str] = None,
|
|
404
|
+
trade_date: Optional[str] = None,
|
|
405
|
+
) -> str:
|
|
406
|
+
"""流式下载原始 Parquet 切片到本地,返回保存路径(消耗下载流量)。
|
|
407
|
+
|
|
408
|
+
trade_date 仅对 tick_data 子分类有效(按交易日下载 {SYM}_{YYYYMMDD}.parquet)。
|
|
409
|
+
"""
|
|
410
|
+
if save_dir is None:
|
|
411
|
+
save_dir = default_download_dir()
|
|
412
|
+
os.makedirs(save_dir, exist_ok=True)
|
|
413
|
+
|
|
414
|
+
params: Dict[str, Any] = {
|
|
415
|
+
"category_id": category_id,
|
|
416
|
+
"sub_category": sub_category,
|
|
417
|
+
}
|
|
418
|
+
if symbol:
|
|
419
|
+
params["symbol"] = symbol
|
|
420
|
+
if trade_date:
|
|
421
|
+
params["trade_date"] = trade_date
|
|
422
|
+
|
|
423
|
+
resp = self._request(
|
|
424
|
+
"GET", "/api/v1/data/download", params=params, stream=True
|
|
425
|
+
)
|
|
426
|
+
if resp.status_code != 200:
|
|
427
|
+
self._check_response(resp)
|
|
428
|
+
|
|
429
|
+
fallback = f"{sub_category}.parquet"
|
|
430
|
+
if symbol:
|
|
431
|
+
fallback = f"{sub_category}_{symbol}.parquet"
|
|
432
|
+
filename = parse_filename_from_content_disposition(
|
|
433
|
+
resp.headers.get("Content-Disposition", ""), fallback
|
|
434
|
+
)
|
|
435
|
+
|
|
436
|
+
save_path = os.path.join(save_dir, filename)
|
|
437
|
+
with open(save_path, "wb") as f:
|
|
438
|
+
for chunk in resp.iter_content(chunk_size=8192):
|
|
439
|
+
if chunk:
|
|
440
|
+
f.write(chunk)
|
|
441
|
+
return os.path.abspath(save_path)
|
|
442
|
+
|
|
443
|
+
def load_as_df(
|
|
444
|
+
self,
|
|
445
|
+
category_id: str,
|
|
446
|
+
sub_category: str,
|
|
447
|
+
symbol: Optional[str] = None,
|
|
448
|
+
trade_date: Optional[str] = None,
|
|
449
|
+
) -> pd.DataFrame:
|
|
450
|
+
"""将远端 Parquet 切片直接加载到内存 DataFrame(不落盘,消耗下载流量)。
|
|
451
|
+
|
|
452
|
+
trade_date 仅对 tick_data 子分类有效(按交易日下载 {SYM}_{YYYYMMDD}.parquet)。
|
|
453
|
+
|
|
454
|
+
带进程内 ETag 缓存:同一对象(ETag 未变)不重复下载,避免对同一 symbol
|
|
455
|
+
多次查询时重复消耗流量。
|
|
456
|
+
"""
|
|
457
|
+
cache_key = f"{category_id}/{sub_category}/{symbol}/{trade_date}"
|
|
458
|
+
params: Dict[str, Any] = {
|
|
459
|
+
"category_id": category_id,
|
|
460
|
+
"sub_category": sub_category,
|
|
461
|
+
}
|
|
462
|
+
if symbol:
|
|
463
|
+
params["symbol"] = symbol
|
|
464
|
+
if trade_date:
|
|
465
|
+
params["trade_date"] = trade_date
|
|
466
|
+
# 若已有缓存,带 If-None-Match 让服务端在对象未变时返回 304(不计流量下载 body)。
|
|
467
|
+
# 注:当前网关透传 COS ETag,但未实现 304;若服务端不支持 304,仍会返回 200
|
|
468
|
+
# 全量 body,此时用响应 ETag 命中本地缓存跳过重复解析。
|
|
469
|
+
cached = self._cache.get(cache_key)
|
|
470
|
+
headers = {}
|
|
471
|
+
if cached and cached.get("etag"):
|
|
472
|
+
headers["If-None-Match"] = cached["etag"]
|
|
473
|
+
resp = self._request("GET", "/api/v1/data/download", params=params, headers=headers)
|
|
474
|
+
if resp.status_code == 304 and cached:
|
|
475
|
+
return cached["df"]
|
|
476
|
+
if resp.status_code != 200:
|
|
477
|
+
self._check_response(resp)
|
|
478
|
+
etag = resp.headers.get("ETag") or ""
|
|
479
|
+
# 命中缓存(ETag 未变):复用已解析的 df,不重复消耗解析
|
|
480
|
+
if cached and etag and cached.get("etag") == etag:
|
|
481
|
+
return cached["df"]
|
|
482
|
+
df = pd.read_parquet(io.BytesIO(resp.content))
|
|
483
|
+
if etag:
|
|
484
|
+
self._cache[cache_key] = {"etag": etag, "df": df}
|
|
485
|
+
return df
|
|
486
|
+
|
|
487
|
+
def query_local(
|
|
488
|
+
self,
|
|
489
|
+
sql: str,
|
|
490
|
+
category_id: str,
|
|
491
|
+
sub_category: str,
|
|
492
|
+
symbol: Optional[str] = None,
|
|
493
|
+
save_dir: Optional[str] = None,
|
|
494
|
+
) -> pd.DataFrame:
|
|
495
|
+
"""用 DuckDB 查询本地已下载的 Parquet 文件(零流量)。
|
|
496
|
+
|
|
497
|
+
若文件不存在会先自动下载。
|
|
498
|
+
|
|
499
|
+
sql 支持两种写法:
|
|
500
|
+
1. 完整 SELECT 语句(必须包含 FROM 子句,表名会被替换为 parquet 路径)
|
|
501
|
+
2. 纯 WHERE 条件字符串(不含 FROM 时自动拼装为 SELECT * FROM <path> WHERE <条件>)
|
|
502
|
+
|
|
503
|
+
安全限制:
|
|
504
|
+
- 不允许分号(;),防止多语句注入
|
|
505
|
+
- 不允许注释(-- 或 /* */),防止注释注入
|
|
506
|
+
- 不允许 UNION / DROP / DELETE / INSERT / UPDATE / ALTER / CREATE / EXEC / EXECUTE
|
|
507
|
+
"""
|
|
508
|
+
file_path = self.download_file(
|
|
509
|
+
category_id=category_id,
|
|
510
|
+
sub_category=sub_category,
|
|
511
|
+
symbol=symbol,
|
|
512
|
+
save_dir=save_dir,
|
|
513
|
+
)
|
|
514
|
+
clean_path = file_path.replace("\\", "/")
|
|
515
|
+
|
|
516
|
+
try:
|
|
517
|
+
import duckdb # type: ignore
|
|
518
|
+
except ImportError as exc:
|
|
519
|
+
raise RuntimeError("本地 SQL 查询需要 duckdb:pip install duckdb") from exc
|
|
520
|
+
|
|
521
|
+
# 安全检查:拒绝危险字符和关键字
|
|
522
|
+
dangerous_keywords = [
|
|
523
|
+
";", "--", "/*", "*/", "union", "drop", "delete",
|
|
524
|
+
"insert", "update", "alter", "create", "exec", "execute",
|
|
525
|
+
]
|
|
526
|
+
sql_upper = sql.upper()
|
|
527
|
+
for kw in dangerous_keywords:
|
|
528
|
+
if kw in sql_upper:
|
|
529
|
+
raise ValidationError(
|
|
530
|
+
f"SQL 包含危险关键字 '{kw}',已被拒绝。"
|
|
531
|
+
"query_local 仅支持 SELECT 查询,不支持数据修改或联合查询。"
|
|
532
|
+
)
|
|
533
|
+
|
|
534
|
+
if "FROM" in sql_upper:
|
|
535
|
+
# 完整 SELECT:替换第一个 FROM 后的表名为 parquet 路径
|
|
536
|
+
sql = re.sub(
|
|
537
|
+
r"FROM\s+[`\"\']?\w+[`\"\']?",
|
|
538
|
+
f"FROM '{clean_path}'",
|
|
539
|
+
sql,
|
|
540
|
+
count=1,
|
|
541
|
+
flags=re.IGNORECASE,
|
|
542
|
+
)
|
|
543
|
+
else:
|
|
544
|
+
# 纯 WHERE 条件
|
|
545
|
+
sql = f"SELECT * FROM '{clean_path}' WHERE {sql}"
|
|
546
|
+
return duckdb.query(sql).df()
|
|
547
|
+
|
|
548
|
+
def get_local_warehouse(self, save_dir: Optional[str] = None) -> "DuckDBWarehouse":
|
|
549
|
+
"""获取 DuckDB 本地数据仓库实例,用于离线多表 SQL JOIN 查询。"""
|
|
550
|
+
return DuckDBWarehouse(client=self, save_dir=save_dir)
|
|
551
|
+
|
|
552
|
+
|
|
553
|
+
class DuckDBWarehouse:
|
|
554
|
+
"""DuckDB 本地仓库管理,支持自动下载数据并注册为 DuckDB 视图,方便多表复杂 SQL 查询。"""
|
|
555
|
+
|
|
556
|
+
def __init__(self, client: QuantDBClient, save_dir: Optional[str] = None):
|
|
557
|
+
try:
|
|
558
|
+
import duckdb # type: ignore
|
|
559
|
+
except ImportError as exc:
|
|
560
|
+
raise RuntimeError("DuckDBWarehouse 需要 duckdb:pip install duckdb") from exc
|
|
561
|
+
|
|
562
|
+
self.client = client
|
|
563
|
+
self.save_dir = save_dir or default_download_dir()
|
|
564
|
+
self.conn = duckdb.connect(database=":memory:")
|
|
565
|
+
self._views: Dict[str, str] = {}
|
|
566
|
+
|
|
567
|
+
def register_file(self, view_name: str, file_path: str) -> None:
|
|
568
|
+
"""把指定的本地 Parquet 文件注册为 DuckDB 视图。"""
|
|
569
|
+
clean_path = file_path.replace("\\", "/")
|
|
570
|
+
self.conn.execute(f"CREATE OR REPLACE VIEW {view_name} AS SELECT * FROM '{clean_path}'")
|
|
571
|
+
self._views[view_name] = clean_path
|
|
572
|
+
|
|
573
|
+
def mount_dataset(
|
|
574
|
+
self,
|
|
575
|
+
category_id: str,
|
|
576
|
+
sub_category: str,
|
|
577
|
+
symbol: Optional[str] = None,
|
|
578
|
+
trade_date: Optional[str] = None,
|
|
579
|
+
view_name: Optional[str] = None,
|
|
580
|
+
) -> str:
|
|
581
|
+
"""自动下载或指定 Parquet 文件,并将其挂载为 DuckDB View。
|
|
582
|
+
|
|
583
|
+
Returns:
|
|
584
|
+
注册的视图名称。
|
|
585
|
+
"""
|
|
586
|
+
file_path = self.client.download_file(
|
|
587
|
+
category_id=category_id,
|
|
588
|
+
sub_category=sub_category,
|
|
589
|
+
symbol=symbol,
|
|
590
|
+
trade_date=trade_date,
|
|
591
|
+
save_dir=self.save_dir,
|
|
592
|
+
)
|
|
593
|
+
if not view_name:
|
|
594
|
+
v_sym = f"_{symbol.replace('.', '_')}" if symbol else ""
|
|
595
|
+
v_td = f"_{trade_date}" if trade_date else ""
|
|
596
|
+
view_name = f"{sub_category}{v_sym}{v_td}"
|
|
597
|
+
# 清理 view_name 中的非法字符
|
|
598
|
+
view_name = re.sub(r"[^a-zA-Z0-9_]", "_", view_name)
|
|
599
|
+
|
|
600
|
+
self.register_file(view_name, file_path)
|
|
601
|
+
return view_name
|
|
602
|
+
|
|
603
|
+
def query(self, sql: str) -> pd.DataFrame:
|
|
604
|
+
"""在所有已注册的视图上执行 SQL 查询,返回 DataFrame。"""
|
|
605
|
+
return self.conn.query(sql).df()
|
|
606
|
+
|
|
607
|
+
def list_views(self) -> Dict[str, str]:
|
|
608
|
+
"""列出当前已注册的视图及其对应本地路径。"""
|
|
609
|
+
return dict(self._views)
|
|
610
|
+
|
quantdb_sdk/errors.py
ADDED
|
@@ -0,0 +1,43 @@
|
|
|
1
|
+
"""QuantDB SDK 异常类型。"""
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class QuantDBError(RuntimeError):
|
|
5
|
+
"""QuantDB SDK 通用异常基类。"""
|
|
6
|
+
|
|
7
|
+
pass
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class AuthError(QuantDBError):
|
|
11
|
+
"""认证失败:API Key 或 JWT Token 无效。"""
|
|
12
|
+
|
|
13
|
+
pass
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class InsufficientTrafficError(QuantDBError):
|
|
17
|
+
"""无有效订阅或本月下载流量已用完。"""
|
|
18
|
+
|
|
19
|
+
pass
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class NotFoundError(QuantDBError):
|
|
23
|
+
"""请求的数据或资源不存在。"""
|
|
24
|
+
|
|
25
|
+
pass
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class RateLimitError(QuantDBError):
|
|
29
|
+
"""请求过于频繁,触发限流。"""
|
|
30
|
+
|
|
31
|
+
pass
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
class ServerError(QuantDBError):
|
|
35
|
+
"""服务端错误。"""
|
|
36
|
+
|
|
37
|
+
pass
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class ValidationError(QuantDBError):
|
|
41
|
+
"""请求参数校验失败。"""
|
|
42
|
+
|
|
43
|
+
pass
|
quantdb_sdk/py.typed
ADDED
|
File without changes
|