"""Load custom data source definitions from user data files.""" from __future__ import annotations import logging import re 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.provider import GenericHTTPProvider logger = logging.getLogger(__name__) _PROVIDERS: dict[str, GenericHTTPProvider] = {} _LOAD_ERRORS: list[dict] = [] _NAME_RE = re.compile(r"^[a-z0-9_]+$") def data_sources_dir() -> Path: return settings.data_dir / "data_sources" def load_all(path: Path | None = None) -> None: """Load all custom provider YAML files into process memory.""" global _PROVIDERS, _LOAD_ERRORS for provider in _PROVIDERS.values(): provider.close() _PROVIDERS = {} _LOAD_ERRORS = [] base = path or data_sources_dir() base.mkdir(parents=True, exist_ok=True) for file in sorted([*base.glob("*.yaml"), *base.glob("*.yml")]): try: config = load_config(file) provider = GenericHTTPProvider(config) errors = provider.validate() if errors: _LOAD_ERRORS.append({"path": str(file), "name": config.name, "errors": errors}) provider.close() continue _PROVIDERS[config.name] = provider except Exception as e: # noqa: BLE001 logger.warning("custom data source load failed %s: %s", file, e) _LOAD_ERRORS.append({"path": str(file), "errors": [str(e)]}) def list_sources() -> list[dict]: return [ { "name": provider.name, "display_name": provider.config.display_name, "datasets": sorted(provider.config.datasets.keys()), "path": str(provider.config.path) if provider.config.path else None, } for provider in _PROVIDERS.values() ] def names() -> set[str]: return set(_PROVIDERS) def errors() -> list[dict]: return list(_LOAD_ERRORS) def get_provider(name: str) -> GenericHTTPProvider: provider = _PROVIDERS.get((name or "").lower()) if provider is None: raise ValueError(f"Custom data source not found or invalid: {name}") return provider def is_custom_provider(name: str) -> bool: return (name or "").lower() in _PROVIDERS def provider_has_dataset(name: str, dataset: str) -> bool: """判断某个 custom 源是否配置了指定数据集。 用于主流程分流: 总开关选了 custom, 但某个数据集未启用时, 该数据集回退 TickFlow。 """ provider = _PROVIDERS.get((name or "").lower()) if provider is None: return False return dataset in provider.config.datasets def get_config_dict(name: str) -> dict | None: """读取一个已加载 custom 源的原始配置 dict(用于前端编辑回填)。""" provider = _PROVIDERS.get((name or "").lower()) if provider is None: return None return _config_to_dict(provider.config) def _config_to_dict(config: CustomSourceConfig) -> dict: auth = config.auth out: dict = { "name": config.name, "display_name": config.display_name, "auth": { "type": auth.type, **({"token_env": auth.token_env} if auth.token_env else {}), **({"header": auth.header} if auth.type in {"bearer", "header"} and auth.header != "Authorization" else {}), **({"param": auth.param} if auth.type == "query" and auth.param != "token" else {}), }, "datasets": {}, } for ds_name, ds in config.datasets.items(): out["datasets"][ds_name] = { "url": ds.url, "method": ds.method, **({"batch": ds.batch} if ds.batch is not None else {}), **({"rpm": ds.rpm} if ds.rpm is not None else {}), "response_path": ds.response_path, "field_map": dict(ds.field_map), **({"transforms": dict(ds.transforms)} if ds.transforms else {}), "symbols_param": ds.symbols_param, "start_param": ds.start_param, "end_param": ds.end_param, } return out def save_config(name: str, config: dict) -> Path: """把一份配置 dict 写成 data/data_sources/{name}.yaml, 返回写入路径。""" if not _NAME_RE.match(name or ""): raise ValueError(f"invalid data source name: {name!r} (only lowercase a-z 0-9 _ allowed)") base = data_sources_dir() base.mkdir(parents=True, exist_ok=True) path = (base / f"{name}.yaml").resolve() if not path.is_relative_to(base.resolve()): raise ValueError("invalid data source name: path escape detected") cleaned = _sanitize_for_yaml(config) path.write_text(yaml.safe_dump(cleaned, allow_unicode=True, sort_keys=False), encoding="utf-8") return path def delete_config(name: str) -> bool: """删除 data/data_sources/{name}.yaml。返回是否真的删除了。""" if not _NAME_RE.match(name or ""): raise ValueError(f"invalid data source name: {name!r}") base = data_sources_dir().resolve() path = (base / f"{name}.yaml").resolve() if not path.is_relative_to(base): raise ValueError("invalid data source name: path escape detected") if not path.exists(): return False path.unlink() return True def _sanitize_for_yaml(config: dict) -> dict: """剔除前端可能塞进来的空值/未启用数据集, 保证写入的 yaml 干净。""" out: dict = { "name": str(config.get("name", "")).lower(), "display_name": str(config.get("display_name") or config.get("name", "")), } auth_raw = config.get("auth") or {} auth_type = str(auth_raw.get("type", "none") or "none").lower() auth: dict = {"type": auth_type} if auth_type != "none" and auth_raw.get("token_env"): auth["token_env"] = str(auth_raw["token_env"]) if auth_type in {"bearer", "header"} and auth_raw.get("header"): auth["header"] = str(auth_raw["header"]) if auth_type == "query" and auth_raw.get("param"): auth["param"] = str(auth_raw["param"]) out["auth"] = auth datasets_out: dict = {} for ds_name, ds_cfg in (config.get("datasets") or {}).items(): if ds_name not in {"daily", "adj_factor", "realtime", "minute", "financial"}: continue if not isinstance(ds_cfg, dict): continue ds = _sanitize_dataset(ds_cfg) if ds: datasets_out[ds_name] = ds out["datasets"] = datasets_out return out def _sanitize_dataset(ds_cfg: dict) -> dict: out: dict = {} url = str(ds_cfg.get("url", "") or "").strip() if not url: return out out["url"] = url method = str(ds_cfg.get("method", "GET") or "GET").upper() out["method"] = method if ds_cfg.get("batch") is not None: try: out["batch"] = int(ds_cfg["batch"]) except (TypeError, ValueError): pass if ds_cfg.get("rpm") is not None: try: out["rpm"] = int(ds_cfg["rpm"]) except (TypeError, ValueError): pass out["response_path"] = str(ds_cfg.get("response_path", "") or "") field_map = { str(k): str(v) for k, v in (ds_cfg.get("field_map") or {}).items() if str(k).strip() and str(v).strip() } if field_map: out["field_map"] = field_map transforms = { str(k): str(v) for k, v in (ds_cfg.get("transforms") or {}).items() if str(k).strip() and str(v).strip() } if transforms: out["transforms"] = transforms if ds_cfg.get("symbols_param"): out["symbols_param"] = str(ds_cfg["symbols_param"]) if ds_cfg.get("start_param"): out["start_param"] = str(ds_cfg["start_param"]) if ds_cfg.get("end_param"): out["end_param"] = str(ds_cfg["end_param"]) return out