sage-api / src /profiles /loader.py
hiasgnpsadgd's picture
Deploy current MIMIC-only SAGE backend
de0f30b verified
Raw
History Blame Contribute Delete
3.1 kB
"""Load a :class:`DataSourceProfile` from ``src/configs/datasets/<name>.yaml``.
The active profile defaults to ``va_cdw`` (today's VA behaviour) and is
overridable via the ``SAGE_DATASOURCE_PROFILE`` environment variable, so the
same pipeline targets MIMIC with ``SAGE_DATASOURCE_PROFILE=mimiciv_duckdb``.
"""
from __future__ import annotations
import functools
import os
from pathlib import Path
import yaml
from .profile import CanonicalView, DataSourceProfile, SourceSpec
# src/profiles/loader.py -> parents[1] == src
SRC_ROOT = Path(__file__).resolve().parents[1]
DATASETS_DIR = SRC_ROOT / "configs" / "datasets"
DEFAULT_PROFILE_NAME = "va_cdw"
ACTIVE_PROFILE_ENV = "SAGE_DATASOURCE_PROFILE"
def _build_source(source_id: str, spec: dict) -> SourceSpec:
return SourceSpec(
source_id=source_id,
path=str(spec["path"]),
source_type=str(spec.get("source_type", "csv")),
meaning=str(spec.get("meaning", "")),
search_columns=tuple(spec.get("search_columns", []) or []),
display_columns=tuple(spec.get("display_columns", []) or []),
)
def load_profile(name: str) -> DataSourceProfile:
path = DATASETS_DIR / f"{name}.yaml"
if not path.exists():
available = sorted(p.stem for p in DATASETS_DIR.glob("*.yaml"))
raise FileNotFoundError(
f"No data-source profile {name!r} at {path}. Available: {available}"
)
raw = yaml.safe_load(path.read_text(encoding="utf-8")) or {}
sources = {
sid: _build_source(sid, spec)
for sid, spec in (raw.get("sources") or {}).items()
}
concept_sources_by_domain = {
domain: tuple(source_ids)
for domain, source_ids in (raw.get("concept_sources_by_domain") or {}).items()
}
domain_concept_keys = {
domain: tuple(keys)
for domain, keys in (raw.get("domain_concept_keys") or {}).items()
}
canonical_views = {
domain: CanonicalView(
domain=domain,
relation=str(cv.get("relation", "")),
roles=dict(cv.get("roles") or {}),
create_view_sql=cv.get("create_view_sql"),
)
for domain, cv in (raw.get("canonical_views") or {}).items()
}
return DataSourceProfile(
name=str(raw.get("name", name)),
dialect=str(raw.get("dialect", "tsql")),
retrieval_mode=str(raw.get("retrieval_mode", "discover")),
sources=sources,
concept_sources_by_domain=concept_sources_by_domain,
domain_concept_keys=domain_concept_keys,
canonical_views=canonical_views,
capabilities=dict(raw.get("capabilities") or {}),
raw=raw,
)
@functools.lru_cache(maxsize=None)
def _load_cached(name: str) -> DataSourceProfile:
return load_profile(name)
def active_profile_name() -> str:
return (os.environ.get(ACTIVE_PROFILE_ENV) or DEFAULT_PROFILE_NAME).strip() or DEFAULT_PROFILE_NAME
def get_active_profile() -> DataSourceProfile:
"""The profile selected by ``$SAGE_DATASOURCE_PROFILE`` (default ``va_cdw``)."""
return _load_cached(active_profile_name())