mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 15:34:16 +08:00
fix(data-source): harden custom source testing
This commit is contained in:
@@ -11,6 +11,7 @@ from fastapi import APIRouter, HTTPException, Request
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app import secrets_store
|
||||
from app.data_providers.custom.config import MAX_TIMEOUT
|
||||
from app.tickflow import client as tf_client
|
||||
from app.tickflow.policy import (
|
||||
detect_capabilities,
|
||||
@@ -349,12 +350,6 @@ class DataProvidersIn(BaseModel):
|
||||
financial_data_provider: str | None = None
|
||||
|
||||
|
||||
class CustomSourceTestIn(BaseModel):
|
||||
provider: str
|
||||
dataset: str
|
||||
symbols: list[str] | None = None
|
||||
|
||||
|
||||
class DatasetFieldMapItem(BaseModel):
|
||||
source: str
|
||||
target: str
|
||||
@@ -373,7 +368,12 @@ class DatasetConfigIn(BaseModel):
|
||||
end_param: str = "end_time"
|
||||
asset_type_param: str | None = None
|
||||
freq_param: str | None = None
|
||||
timeout: float | None = Field(default=None, gt=0, allow_inf_nan=False)
|
||||
timeout: float | None = Field(
|
||||
default=None,
|
||||
gt=0,
|
||||
le=MAX_TIMEOUT,
|
||||
allow_inf_nan=False,
|
||||
)
|
||||
|
||||
|
||||
class AuthConfigIn(BaseModel):
|
||||
@@ -390,6 +390,13 @@ class CustomSourceIn(BaseModel):
|
||||
datasets: dict[str, DatasetConfigIn] = {}
|
||||
|
||||
|
||||
class CustomSourceTestIn(BaseModel):
|
||||
provider: str
|
||||
dataset: str
|
||||
symbols: list[str] | None = None
|
||||
config: CustomSourceIn | None = None
|
||||
|
||||
|
||||
@router.get("/preferences")
|
||||
def get_preferences() -> dict:
|
||||
"""返回用户偏好设置。"""
|
||||
@@ -569,11 +576,25 @@ def delete_data_source(name: str) -> dict:
|
||||
def test_data_source(req: CustomSourceTestIn) -> dict:
|
||||
"""试拉自定义数据源,不写盘。"""
|
||||
from app.data_providers import custom as custom_sources
|
||||
provider = custom_sources.get_provider(req.provider)
|
||||
|
||||
temporary = req.config is not None
|
||||
provider = None
|
||||
try:
|
||||
if req.config:
|
||||
config = req.config.model_dump()
|
||||
dataset_config = config["datasets"].get(req.dataset)
|
||||
if dataset_config is None:
|
||||
raise ValueError(f"dataset '{req.dataset}' is not configured")
|
||||
config["datasets"] = {req.dataset: dataset_config}
|
||||
provider = custom_sources.create_provider(config)
|
||||
else:
|
||||
provider = custom_sources.get_provider(req.provider)
|
||||
return provider.test_dataset(req.dataset, req.symbols)
|
||||
except Exception as e: # noqa: BLE001
|
||||
raise HTTPException(status_code=400, detail=f"自定义数据源测试失败: {e}") from e
|
||||
finally:
|
||||
if temporary and provider is not None:
|
||||
provider.close()
|
||||
|
||||
|
||||
@router.put("/preferences/data-providers")
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
"""Custom data source extension points."""
|
||||
from app.data_providers.custom.loader import (
|
||||
create_provider,
|
||||
data_sources_dir,
|
||||
delete_config,
|
||||
errors,
|
||||
@@ -18,6 +19,7 @@ from app.data_providers.custom.loader import (
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"create_provider",
|
||||
"data_sources_dir",
|
||||
"delete_config",
|
||||
"errors",
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
"""Custom HTTP data source configuration."""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
@@ -9,6 +8,8 @@ from typing import Any, Literal
|
||||
import yaml
|
||||
|
||||
DatasetName = Literal["daily", "adj_factor", "realtime", "minute", "financial"]
|
||||
DEFAULT_TIMEOUT = 30.0
|
||||
MAX_TIMEOUT = 300.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -25,7 +26,7 @@ class DatasetConfig:
|
||||
method: str = "GET"
|
||||
batch: int | None = None
|
||||
rpm: int | None = None
|
||||
timeout: float = 30.0
|
||||
timeout: float = DEFAULT_TIMEOUT
|
||||
response_path: str = ""
|
||||
field_map: dict[str, str] = field(default_factory=dict)
|
||||
transforms: dict[str, str] = field(default_factory=dict)
|
||||
@@ -61,12 +62,18 @@ def _auth_from_dict(raw: dict[str, Any] | None) -> AuthConfig:
|
||||
|
||||
|
||||
def _dataset_from_dict(raw: dict[str, Any]) -> DatasetConfig:
|
||||
try:
|
||||
timeout = float(raw.get("timeout", 30.0) or 30.0)
|
||||
except (TypeError, ValueError):
|
||||
timeout = 30.0
|
||||
if not math.isfinite(timeout) or timeout <= 0:
|
||||
timeout = 30.0
|
||||
timeout_raw = raw.get("timeout")
|
||||
if timeout_raw is None:
|
||||
timeout = DEFAULT_TIMEOUT
|
||||
else:
|
||||
try:
|
||||
timeout = float(timeout_raw)
|
||||
except (TypeError, ValueError) as e:
|
||||
raise ValueError(
|
||||
f"timeout must be a number between 0 and {MAX_TIMEOUT:g} seconds"
|
||||
) from e
|
||||
if not 0 < timeout <= MAX_TIMEOUT:
|
||||
raise ValueError(f"timeout must be between 0 and {MAX_TIMEOUT:g} seconds")
|
||||
|
||||
return DatasetConfig(
|
||||
url=str(raw.get("url", "") or ""),
|
||||
@@ -87,14 +94,14 @@ def _dataset_from_dict(raw: dict[str, Any]) -> DatasetConfig:
|
||||
)
|
||||
|
||||
|
||||
def load_config(path: Path) -> CustomSourceConfig:
|
||||
raw = yaml.safe_load(path.read_text(encoding="utf-8")) or {}
|
||||
def config_from_dict(raw: dict[str, Any], path: Path | None = None) -> CustomSourceConfig:
|
||||
datasets = {
|
||||
name: _dataset_from_dict(cfg)
|
||||
for name, cfg in (raw.get("datasets") or {}).items()
|
||||
if name in {"daily", "adj_factor", "realtime", "minute", "financial"} and isinstance(cfg, dict)
|
||||
}
|
||||
name = str(raw.get("name", path.stem) or path.stem).lower()
|
||||
default_name = path.stem if path else "preview"
|
||||
name = str(raw.get("name", default_name) or default_name).lower()
|
||||
return CustomSourceConfig(
|
||||
name=name,
|
||||
display_name=str(raw.get("display_name", name) or name),
|
||||
@@ -102,3 +109,8 @@ def load_config(path: Path) -> CustomSourceConfig:
|
||||
datasets=datasets,
|
||||
path=path,
|
||||
)
|
||||
|
||||
|
||||
def load_config(path: Path) -> CustomSourceConfig:
|
||||
raw = yaml.safe_load(path.read_text(encoding="utf-8")) or {}
|
||||
return config_from_dict(raw, path)
|
||||
|
||||
@@ -12,7 +12,13 @@ from pathlib import Path
|
||||
import yaml
|
||||
|
||||
from app.config import settings
|
||||
from app.data_providers.custom.config import CustomSourceConfig, load_config
|
||||
from app.data_providers.custom.config import (
|
||||
DEFAULT_TIMEOUT,
|
||||
MAX_TIMEOUT,
|
||||
CustomSourceConfig,
|
||||
config_from_dict,
|
||||
load_config,
|
||||
)
|
||||
from app.data_providers.custom.provider import GenericHTTPProvider
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -247,6 +253,17 @@ def get_provider(name: str) -> GenericHTTPProvider:
|
||||
return provider
|
||||
|
||||
|
||||
def create_provider(config: dict) -> GenericHTTPProvider:
|
||||
"""Build a validated temporary provider without writing or registering it."""
|
||||
parsed = config_from_dict(_sanitize_for_yaml(config))
|
||||
provider = GenericHTTPProvider(parsed)
|
||||
errors = provider.validate()
|
||||
if errors:
|
||||
provider.close()
|
||||
raise ValueError("; ".join(errors))
|
||||
return provider
|
||||
|
||||
|
||||
def is_custom_provider(name: str) -> bool:
|
||||
return (name or "").lower() in _PROVIDERS
|
||||
|
||||
@@ -291,7 +308,7 @@ def _config_to_dict(config: CustomSourceConfig) -> dict:
|
||||
"method": ds.method,
|
||||
**({"batch": ds.batch} if ds.batch is not None else {}),
|
||||
**({"rpm": ds.rpm} if ds.rpm is not None else {}),
|
||||
**({"timeout": ds.timeout} if ds.timeout != 30.0 else {}),
|
||||
**({"timeout": ds.timeout} if ds.timeout != DEFAULT_TIMEOUT else {}),
|
||||
"response_path": ds.response_path,
|
||||
"field_map": dict(ds.field_map),
|
||||
**({"transforms": dict(ds.transforms)} if ds.transforms else {}),
|
||||
@@ -389,10 +406,17 @@ def _sanitize_dataset(ds_name: str, ds_cfg: dict) -> dict:
|
||||
if ds_cfg.get("timeout") is not None:
|
||||
try:
|
||||
timeout = float(ds_cfg["timeout"])
|
||||
if math.isfinite(timeout) and timeout > 0:
|
||||
out["timeout"] = timeout
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
except (TypeError, ValueError) as e:
|
||||
raise ValueError(
|
||||
f"{ds_name}: timeout must be a number between 0 and "
|
||||
f"{MAX_TIMEOUT:g} seconds"
|
||||
) from e
|
||||
if not math.isfinite(timeout) or not 0 < timeout <= MAX_TIMEOUT:
|
||||
raise ValueError(
|
||||
f"{ds_name}: timeout must be between 0 and {MAX_TIMEOUT:g} seconds"
|
||||
)
|
||||
if timeout != DEFAULT_TIMEOUT:
|
||||
out["timeout"] = timeout
|
||||
out["response_path"] = str(ds_cfg.get("response_path", "") or "")
|
||||
field_map = {
|
||||
str(k): str(v)
|
||||
@@ -425,6 +449,22 @@ def _sanitize_dataset(ds_name: str, ds_cfg: dict) -> dict:
|
||||
out["asset_type_param"] = asset_type_param
|
||||
if freq_param:
|
||||
out["freq_param"] = freq_param
|
||||
request_params = [
|
||||
out.get("symbols_param", "symbols"),
|
||||
out.get("start_param", "start_time"),
|
||||
out.get("end_param", "end_time"),
|
||||
]
|
||||
if ds_name == "minute":
|
||||
request_params.extend(
|
||||
name for name in (out.get("asset_type_param"), out.get("freq_param")) if name
|
||||
)
|
||||
duplicates = sorted({
|
||||
name for name in request_params if request_params.count(name) > 1
|
||||
})
|
||||
if ds_name != "realtime" and duplicates:
|
||||
raise ValueError(
|
||||
f"{ds_name}: duplicate request parameter names: {', '.join(duplicates)}"
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
import logging
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from datetime import datetime
|
||||
from datetime import datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
@@ -14,7 +14,12 @@ import polars as pl
|
||||
from app.config import settings
|
||||
from app.data_providers.base import AssetType
|
||||
from app.data_providers.custom.config import CustomSourceConfig, DatasetConfig
|
||||
from app.data_providers.custom.mapper import apply_transforms, datetime_payload, extract_rows, map_rows
|
||||
from app.data_providers.custom.mapper import (
|
||||
apply_transforms,
|
||||
datetime_payload,
|
||||
extract_rows,
|
||||
map_rows,
|
||||
)
|
||||
from app.data_providers.normalizer import normalize_adj_factors, normalize_daily
|
||||
from app.tickflow.rate_limits import chunked, sleep_between_batches
|
||||
|
||||
@@ -52,6 +57,20 @@ class GenericHTTPProvider:
|
||||
missing = sorted(required - mapped)
|
||||
if missing:
|
||||
errors.append(f"{dataset}: missing mapped fields: {', '.join(missing)}")
|
||||
if dataset != "realtime":
|
||||
request_params = [cfg.symbols_param, cfg.start_param, cfg.end_param]
|
||||
if dataset == "minute":
|
||||
request_params.extend(
|
||||
name for name in (cfg.asset_type_param, cfg.freq_param) if name
|
||||
)
|
||||
duplicates = sorted({
|
||||
name for name in request_params if request_params.count(name) > 1
|
||||
})
|
||||
if duplicates:
|
||||
errors.append(
|
||||
f"{dataset}: duplicate request parameter names: "
|
||||
f"{', '.join(duplicates)}"
|
||||
)
|
||||
return errors
|
||||
|
||||
def get_daily(
|
||||
@@ -192,7 +211,34 @@ class GenericHTTPProvider:
|
||||
|
||||
def test_dataset(self, dataset: str, symbols: list[str] | None = None) -> dict:
|
||||
cfg = self._dataset(dataset)
|
||||
rows = self._request_rows(cfg, symbols=symbols or [])
|
||||
test_symbols = symbols or ["000001.SZ"]
|
||||
end_time = datetime.now()
|
||||
start_time = end_time - timedelta(days=7)
|
||||
if dataset == "realtime":
|
||||
rows = self._request_rows(cfg)
|
||||
elif dataset == "minute":
|
||||
override: dict[str, Any] = {}
|
||||
if cfg.asset_type_param:
|
||||
override[cfg.asset_type_param] = "stock"
|
||||
if cfg.freq_param:
|
||||
override[cfg.freq_param] = "1m"
|
||||
rows = self._request_rows(
|
||||
cfg,
|
||||
symbols=test_symbols,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
override_params=override or None,
|
||||
override_body=override or None,
|
||||
)
|
||||
elif dataset in {"daily", "adj_factor"}:
|
||||
rows = self._request_rows(
|
||||
cfg,
|
||||
symbols=test_symbols,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
else:
|
||||
rows = self._request_rows(cfg, symbols=test_symbols)
|
||||
df = self._mapped_frame(cfg, rows)
|
||||
return {
|
||||
"provider": self.name,
|
||||
|
||||
@@ -1,11 +1,29 @@
|
||||
import math
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from pydantic import ValidationError
|
||||
|
||||
from app.api.settings import DatasetConfigIn
|
||||
from app.data_providers.custom.config import CustomSourceConfig, _dataset_from_dict
|
||||
from app.api.settings import (
|
||||
CustomSourceIn,
|
||||
CustomSourceTestIn,
|
||||
DatasetConfigIn,
|
||||
)
|
||||
from app.api.settings import (
|
||||
test_data_source as run_data_source_test,
|
||||
)
|
||||
from app.data_providers.custom.config import (
|
||||
MAX_TIMEOUT,
|
||||
CustomSourceConfig,
|
||||
DatasetConfig,
|
||||
_dataset_from_dict,
|
||||
load_config,
|
||||
)
|
||||
from app.data_providers.custom.loader import _config_to_dict, _sanitize_for_yaml
|
||||
from app.data_providers.custom.provider import GenericHTTPProvider
|
||||
|
||||
|
||||
def test_minute_request_parameter_names_survive_config_round_trip():
|
||||
@@ -77,30 +95,67 @@ def test_timeout_survives_config_round_trip():
|
||||
assert "start_param" not in realtime
|
||||
assert "end_param" not in realtime
|
||||
|
||||
explicit_default = _sanitize_for_yaml({
|
||||
"name": "test_source",
|
||||
"datasets": {
|
||||
"daily": {
|
||||
"url": "https://example.test/daily",
|
||||
"timeout": 30.0,
|
||||
},
|
||||
},
|
||||
})
|
||||
assert "timeout" not in explicit_default["datasets"]["daily"]
|
||||
|
||||
@pytest.mark.parametrize("timeout", [0, -1, math.nan, math.inf, -math.inf])
|
||||
def test_timeout_api_rejects_non_positive_or_non_finite_values(timeout):
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"timeout",
|
||||
[0, -1, math.nan, math.inf, -math.inf, MAX_TIMEOUT + 1],
|
||||
)
|
||||
def test_timeout_api_rejects_out_of_range_or_non_finite_values(timeout):
|
||||
with pytest.raises(ValidationError):
|
||||
DatasetConfigIn(url="https://example.test/daily", timeout=timeout)
|
||||
|
||||
|
||||
def test_invalid_yaml_timeout_falls_back_to_default():
|
||||
for timeout in (0, -1, math.nan, math.inf, -math.inf, "invalid"):
|
||||
parsed = _dataset_from_dict({
|
||||
"url": "https://example.test/daily",
|
||||
"timeout": timeout,
|
||||
})
|
||||
assert parsed.timeout == 30.0
|
||||
def test_invalid_yaml_timeout_is_a_load_error(tmp_path: Path):
|
||||
for index, timeout in enumerate((0, -1, math.nan, math.inf, -math.inf, "invalid")):
|
||||
path = tmp_path / f"invalid_{index}.yaml"
|
||||
path.write_text(
|
||||
"\n".join([
|
||||
"name: invalid",
|
||||
"datasets:",
|
||||
" daily:",
|
||||
" url: https://example.test/daily",
|
||||
f" timeout: {timeout}",
|
||||
]),
|
||||
encoding="utf-8",
|
||||
)
|
||||
with pytest.raises(ValueError, match="timeout must be"):
|
||||
load_config(path)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("timeout", [0, -1, math.nan, math.inf, -math.inf, "invalid"])
|
||||
def test_sanitizer_drops_invalid_timeout_and_realtime_request_params(timeout):
|
||||
def test_sanitizer_rejects_invalid_timeout(timeout):
|
||||
with pytest.raises(ValueError, match="timeout must be"):
|
||||
_sanitize_for_yaml({
|
||||
"name": "test_source",
|
||||
"datasets": {
|
||||
"realtime": {
|
||||
"url": "https://example.test/realtime",
|
||||
"timeout": timeout,
|
||||
"symbols_param": "codes",
|
||||
"start_param": "from",
|
||||
"end_param": "to",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
|
||||
def test_sanitizer_drops_realtime_request_params():
|
||||
cleaned = _sanitize_for_yaml({
|
||||
"name": "test_source",
|
||||
"datasets": {
|
||||
"realtime": {
|
||||
"url": "https://example.test/realtime",
|
||||
"timeout": timeout,
|
||||
"symbols_param": "codes",
|
||||
"start_param": "from",
|
||||
"end_param": "to",
|
||||
@@ -109,7 +164,6 @@ def test_sanitizer_drops_invalid_timeout_and_realtime_request_params(timeout):
|
||||
})
|
||||
|
||||
dataset = cleaned["datasets"]["realtime"]
|
||||
assert "timeout" not in dataset
|
||||
assert "symbols_param" not in dataset
|
||||
assert "start_param" not in dataset
|
||||
assert "end_param" not in dataset
|
||||
@@ -132,3 +186,141 @@ def test_empty_request_parameter_names_restore_defaults():
|
||||
assert parsed.symbols_param == "symbols"
|
||||
assert parsed.start_param == "start_time"
|
||||
assert parsed.end_param == "end_time"
|
||||
|
||||
|
||||
def _dataset_config(dataset: str = "minute", **overrides) -> DatasetConfig:
|
||||
required = {
|
||||
"daily": ("symbol", "date", "open", "high", "low", "close", "volume", "amount"),
|
||||
"adj_factor": ("symbol", "trade_date", "ex_factor"),
|
||||
"realtime": ("symbol", "last_price", "prev_close", "open", "high", "low", "volume"),
|
||||
"minute": ("symbol", "datetime", "open", "high", "low", "close", "volume", "amount"),
|
||||
}
|
||||
values = {
|
||||
"url": f"https://example.test/{dataset}",
|
||||
"field_map": {name: name for name in required[dataset]},
|
||||
**overrides,
|
||||
}
|
||||
return DatasetConfig(**values)
|
||||
|
||||
|
||||
def _capture_test_request(dataset: str, **overrides):
|
||||
provider = GenericHTTPProvider(CustomSourceConfig(
|
||||
name="test_source",
|
||||
display_name="Test Source",
|
||||
datasets={dataset: _dataset_config(dataset, **overrides)},
|
||||
))
|
||||
captured = {}
|
||||
|
||||
def request_rows(cfg, **kwargs):
|
||||
captured.update(kwargs)
|
||||
return []
|
||||
|
||||
provider._request_rows = request_rows
|
||||
provider.test_dataset(dataset, ["600000.SH"])
|
||||
provider.close()
|
||||
return captured
|
||||
|
||||
|
||||
def test_test_dataset_realtime_omits_symbol_and_time_parameters():
|
||||
assert _capture_test_request("realtime") == {}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dataset", ["daily", "adj_factor"])
|
||||
def test_test_dataset_history_uses_symbol_and_short_time_range(dataset):
|
||||
captured = _capture_test_request(dataset)
|
||||
|
||||
assert captured["symbols"] == ["600000.SH"]
|
||||
assert isinstance(captured["start_time"], datetime)
|
||||
assert isinstance(captured["end_time"], datetime)
|
||||
assert (captured["end_time"] - captured["start_time"]).days == 7
|
||||
|
||||
|
||||
def test_test_dataset_minute_injects_production_overrides():
|
||||
captured = _capture_test_request(
|
||||
"minute",
|
||||
asset_type_param="asset",
|
||||
freq_param="period",
|
||||
)
|
||||
|
||||
assert captured["symbols"] == ["600000.SH"]
|
||||
assert captured["override_params"] == {"asset": "stock", "period": "1m"}
|
||||
assert captured["override_body"] == {"asset": "stock", "period": "1m"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dataset", ["daily", "minute"])
|
||||
def test_duplicate_dynamic_parameter_names_are_rejected(dataset):
|
||||
config = _dataset_config(dataset, symbols_param="range", start_param="range")
|
||||
provider = GenericHTTPProvider(CustomSourceConfig(
|
||||
name="test_source",
|
||||
display_name="Test Source",
|
||||
datasets={dataset: config},
|
||||
))
|
||||
|
||||
try:
|
||||
assert provider.validate() == [
|
||||
f"{dataset}: duplicate request parameter names: range"
|
||||
]
|
||||
finally:
|
||||
provider.close()
|
||||
|
||||
with pytest.raises(ValueError, match="duplicate request parameter names: range"):
|
||||
_sanitize_for_yaml({
|
||||
"name": "test_source",
|
||||
"datasets": {
|
||||
dataset: {
|
||||
"url": config.url,
|
||||
"symbols_param": "range",
|
||||
"start_param": "range",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
|
||||
def test_data_source_trial_uses_unsaved_config_and_closes_provider(monkeypatch):
|
||||
from app.data_providers import custom as custom_sources
|
||||
|
||||
provider = Mock()
|
||||
provider.test_dataset.return_value = {
|
||||
"provider": "draft",
|
||||
"dataset": "realtime",
|
||||
"rows": 0,
|
||||
"columns": [],
|
||||
"preview": [],
|
||||
}
|
||||
create_provider = Mock(return_value=provider)
|
||||
monkeypatch.setattr(custom_sources, "create_provider", create_provider)
|
||||
config = CustomSourceIn(
|
||||
name="draft",
|
||||
datasets={
|
||||
"daily": DatasetConfigIn(url="https://unfinished.test"),
|
||||
"realtime": DatasetConfigIn(url="https://example.test/realtime"),
|
||||
},
|
||||
)
|
||||
|
||||
result = run_data_source_test(CustomSourceTestIn(
|
||||
provider="draft",
|
||||
dataset="realtime",
|
||||
config=config,
|
||||
))
|
||||
|
||||
assert result["provider"] == "draft"
|
||||
tested = create_provider.call_args.args[0]
|
||||
assert list(tested["datasets"]) == ["realtime"]
|
||||
provider.test_dataset.assert_called_once_with("realtime", None)
|
||||
provider.close.assert_called_once_with()
|
||||
|
||||
|
||||
def test_data_source_trial_wraps_missing_saved_provider_as_http_400(monkeypatch):
|
||||
from app.data_providers import custom as custom_sources
|
||||
|
||||
monkeypatch.setattr(
|
||||
custom_sources,
|
||||
"get_provider",
|
||||
Mock(side_effect=ValueError("not found")),
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
run_data_source_test(CustomSourceTestIn(provider="missing", dataset="daily"))
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "not found" in exc_info.value.detail
|
||||
|
||||
@@ -157,7 +157,7 @@ freq_param: period
|
||||
timeout: 60
|
||||
```
|
||||
|
||||
留空或省略时用默认 30 秒;该值对数据同步与「试拉测试」均生效。在设置页编辑数据源时可在「超时」输入框修改(与 批量 / RPM / 响应路径 同行)。
|
||||
留空或省略时用默认 30 秒,可配置范围为大于 0 且不超过 300 秒;该值对数据同步与「试拉测试」均生效。在设置页编辑数据源时可在「超时」输入框修改(与 批量 / RPM / 响应路径 同行)。「试拉测试」直接使用当前表单内容,新建数据源或尚未保存的修改也可测试。
|
||||
|
||||
## 鉴权
|
||||
|
||||
|
||||
@@ -1018,10 +1018,15 @@ export const api = {
|
||||
`/api/settings/plugins/${encodeURIComponent(name)}/install`,
|
||||
{ method: 'DELETE' },
|
||||
),
|
||||
testDataSource: (provider: string, dataset: string, symbols?: string[]) =>
|
||||
testDataSource: (
|
||||
provider: string,
|
||||
dataset: string,
|
||||
symbols?: string[],
|
||||
config?: CustomSourceConfig,
|
||||
) =>
|
||||
request<DataSourceTestResult>('/api/settings/data-sources/test', {
|
||||
method: 'POST',
|
||||
body: JSON.stringify({ provider, dataset, symbols }),
|
||||
body: JSON.stringify({ provider, dataset, symbols, config }),
|
||||
}),
|
||||
updateDataProviders: (cfg: Partial<Pick<Preferences, 'daily_data_provider' | 'adj_factor_provider' | 'minute_data_provider' | 'realtime_data_provider' | 'financial_data_provider'>>) =>
|
||||
request<Pick<Preferences, 'daily_data_provider' | 'adj_factor_provider' | 'minute_data_provider' | 'realtime_data_provider'>>(
|
||||
|
||||
@@ -50,6 +50,33 @@ const FIELD_LABELS: Record<string, string> = {
|
||||
session: '交易时段',
|
||||
}
|
||||
|
||||
function normalizeConfig(config: CustomSourceConfig): CustomSourceConfig {
|
||||
return {
|
||||
...config,
|
||||
name: config.name.toLowerCase().trim(),
|
||||
display_name: config.display_name.trim() || config.name.toLowerCase().trim(),
|
||||
datasets: Object.fromEntries(
|
||||
Object.entries(config.datasets).map(([key, dataset]) => {
|
||||
const normalized = { ...dataset }
|
||||
if (key === 'realtime') {
|
||||
delete normalized.symbols_param
|
||||
delete normalized.start_param
|
||||
delete normalized.end_param
|
||||
}
|
||||
return [
|
||||
key,
|
||||
{
|
||||
...normalized,
|
||||
field_map: Object.fromEntries(
|
||||
Object.entries(dataset.field_map).filter(([source]) => !source.startsWith('__pending_'))
|
||||
),
|
||||
},
|
||||
]
|
||||
})
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
function emptyConfig(): CustomSourceConfig {
|
||||
return { name: '', display_name: '', auth: { type: 'none' }, datasets: {} }
|
||||
}
|
||||
@@ -96,36 +123,16 @@ export function DataSourceEditor({
|
||||
if (!ds.url.trim()) {
|
||||
throw new Error(`数据集「${DATASET_LABEL[key as DatasetKey] || key}」未填写接口 URL`)
|
||||
}
|
||||
if (ds.timeout != null && (!Number.isFinite(ds.timeout) || ds.timeout <= 0)) {
|
||||
throw new Error(`数据集「${DATASET_LABEL[key as DatasetKey] || key}」超时必须大于 0 秒`)
|
||||
if (
|
||||
ds.timeout != null &&
|
||||
(!Number.isFinite(ds.timeout) || ds.timeout <= 0 || ds.timeout > 300)
|
||||
) {
|
||||
throw new Error(
|
||||
`数据集「${DATASET_LABEL[key as DatasetKey] || key}」超时必须在 0 到 300 秒之间`
|
||||
)
|
||||
}
|
||||
}
|
||||
// 提交时去掉 field_map 里的 __pending_ 临时 key (未填外部字段名的草稿行)
|
||||
const cleaned: CustomSourceConfig = {
|
||||
...config,
|
||||
name: config.name.toLowerCase().trim(),
|
||||
display_name: config.display_name.trim() || config.name.toLowerCase().trim(),
|
||||
datasets: Object.fromEntries(
|
||||
Object.entries(config.datasets).map(([k, ds]) => {
|
||||
const normalized = { ...ds }
|
||||
if (k === 'realtime') {
|
||||
delete normalized.symbols_param
|
||||
delete normalized.start_param
|
||||
delete normalized.end_param
|
||||
}
|
||||
return [
|
||||
k,
|
||||
{
|
||||
...normalized,
|
||||
field_map: Object.fromEntries(
|
||||
Object.entries(ds.field_map).filter(([src]) => !src.startsWith('__pending_'))
|
||||
),
|
||||
},
|
||||
]
|
||||
})
|
||||
),
|
||||
}
|
||||
return api.saveDataSource(cleaned)
|
||||
return api.saveDataSource(normalizeConfig(config))
|
||||
},
|
||||
onSuccess: () => {
|
||||
toast(isNew ? '数据源已创建' : '数据源已更新', 'success')
|
||||
@@ -279,6 +286,7 @@ export function DataSourceEditor({
|
||||
<div className="p-5">
|
||||
<DatasetDetail
|
||||
key={activeTab}
|
||||
config={config}
|
||||
datasetKey={activeTab}
|
||||
cfg={config.datasets[activeTab]}
|
||||
providerName={config.name.toLowerCase().trim() || existingName || ''}
|
||||
@@ -314,6 +322,7 @@ export function DataSourceEditor({
|
||||
}
|
||||
|
||||
function DatasetDetail({
|
||||
config,
|
||||
datasetKey,
|
||||
cfg,
|
||||
providerName,
|
||||
@@ -321,6 +330,7 @@ function DatasetDetail({
|
||||
onFieldMap,
|
||||
onToggle,
|
||||
}: {
|
||||
config: CustomSourceConfig
|
||||
datasetKey: DatasetKey
|
||||
cfg?: DatasetConfig
|
||||
providerName: string
|
||||
@@ -337,6 +347,10 @@ function DatasetDetail({
|
||||
providerName,
|
||||
datasetKey,
|
||||
testSymbols.split(/[,\s]+/).map(s => s.trim()).filter(Boolean),
|
||||
normalizeConfig({
|
||||
...config,
|
||||
datasets: cfg ? { [datasetKey]: cfg } : {},
|
||||
}),
|
||||
),
|
||||
})
|
||||
|
||||
@@ -402,6 +416,7 @@ function DatasetDetail({
|
||||
<input
|
||||
type="number"
|
||||
min="0.1"
|
||||
max="300"
|
||||
step="any"
|
||||
value={cfg.timeout ?? ''}
|
||||
onChange={e => onUpdate({ timeout: e.target.value ? Number(e.target.value) : null })}
|
||||
@@ -539,7 +554,7 @@ function DatasetDetail({
|
||||
/>
|
||||
<button
|
||||
onClick={() => test.mutate()}
|
||||
disabled={test.isPending || !cfg.url || !providerName}
|
||||
disabled={test.isPending || !cfg.url}
|
||||
className="inline-flex items-center gap-1 px-3 py-1.5 rounded-btn bg-elevated text-secondary hover:text-foreground text-xs disabled:opacity-40 transition-colors"
|
||||
>
|
||||
{test.isPending ? '测试中...' : '测试'}
|
||||
|
||||
Reference in New Issue
Block a user