diff --git a/CHANGELOG.md b/CHANGELOG.md index efcc8ad..1c8d34e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,20 @@ ## Unreleased +- Added PIT release rules (`alphaforge.pit.release_rules`): `NthBusinessDay`, `NthWeekday`, `CalendarDay`, `FixedLagMonths`, `QuarterlyRelease`, `WeeklyRelease`, and `CustomRule` with a tagged-union registry for YAML round-trips. +- Added vintage resolvers (`alphaforge.pit.resolvers`): `RealtimeResolver`, `LatestResolver`, and `FrozenResolver` implementing the `VintageResolver` protocol for point-in-time backtesting views. +- Added `VintageView` value object (`alphaforge.pit.views`) to declare realtime / latest / frozen vintage strategies without coupling to resolution logic. +- Added PIT panel builder utilities (`alphaforge.pit.panel`): `build_pit_panel` and `long_to_wide` for assembling aligned panels from PIT snapshots. +- Added missingness taxonomy (`alphaforge.pit.missingness`) for classifying NaN cells in nowcasting panels by cause. +- Added vintage selection and lookahead validation utilities (`alphaforge.pit.vintage`): `select_vintage_for_asof` and `validate_no_lookahead`. +- Fixed CI linting errors: removed unused imports and sorted import blocks across new modules and test files. +- Fixed mypy type errors: tightened `_register` signature in `pit/release_rules.py`, added `name: str` to `Parametric` protocol in `pipeline/protocols.py`, and corrected `Mapping[Any, str]` annotation in `data/short_rates.py`. + +- Added `alphaforge.evaluation` package with pluggable metric infrastructure: + - `MetricFn` protocol (runtime-checkable) for composable forecast accuracy metrics. + - Built-in implementations: `RMSE`, `MAE`, `DirectionalAccuracy`, `MAPE`, `MeanError`. + - Pre-built suites: `DEFAULT_METRICS` (RMSE + MAE + DA), `BENCHMARK_METRICS` (adds MeanError + MAPE). + - API reference page: `docs/api/evaluation-metrics.md`. - Fixed type annotation errors in `pit/accessor.py` and `pit/models.py` (stale `type: ignore` comments, TypedDict narrowing, and `ast.Call` attribute access error codes). - Fixed linting errors: sorted import blocks (ruff I001) in `alphaforge/__init__.py`, `alphaforge/pit/__init__.py`, `alphaforge/pit/accessor.py`, `alphaforge/pit/gdp.py`, `alphaforge/pit/tasks.py`; removed unused imports (ruff F401) in `alphaforge/data/public_web/cftc_cot.py`. - Fixed mypy type narrowing error in `iter_walk_forward_folds` (`pit/tasks.py`): replaced `int(min_train_size)` with a direct reference guarded by a type-narrowing assertion. diff --git a/alphaforge/__init__.py b/alphaforge/__init__.py index 9b2bbc6..394f82f 100644 --- a/alphaforge/__init__.py +++ b/alphaforge/__init__.py @@ -4,13 +4,6 @@ from .data.fred_source import FREDDataSource from .data.panel import PanelFrame from .data.pit_source import PITDataSource -from .data.short_rates import ( - ShortRateDataset, - build_duan_weekly_dataset, - build_kim_orphanides_dataset, - build_macro_finance_dataset, - build_policy_rule_dataset, -) from .data.public_web import ( ANPFuelPricesDataSource, B3HistoricalQuotesDataSource, @@ -35,6 +28,13 @@ ) from .data.query import Query from .data.schema import TableSchema +from .data.short_rates import ( + ShortRateDataset, + build_duan_weekly_dataset, + build_kim_orphanides_dataset, + build_macro_finance_dataset, + build_policy_rule_dataset, +) from .data.universe import EntityMetadata, Universe from .features.frame import Artifact, FeatureFrame from .features.ops import join_feature_frames, materialize @@ -92,6 +92,7 @@ ) from .pit.transforms import PITTransformResult, PITTransformSpec from .pit.validation import PITValidationReport, validate_pit_observations +from .registry import EntityEntry, EntityRegistry from .store.cache import MaterializationPolicy from .store.duckdb_parquet import DuckDBParquetStore from .time.align import AlignedPanel, AlignSpec, AvailabilityState, align_panel @@ -206,4 +207,6 @@ "build_snapshot_tape", "make_ref_entity_id", "parse_ref_entity_id", + "EntityEntry", + "EntityRegistry", ] diff --git a/alphaforge/config.py b/alphaforge/config.py new file mode 100644 index 0000000..074e8d8 --- /dev/null +++ b/alphaforge/config.py @@ -0,0 +1,27 @@ +"""Generic configuration resolution utilities.""" +from __future__ import annotations + +import os +from dataclasses import dataclass + + +def resolve_config_value( + explicit: str | None, + env_var: str, + default: str, +) -> str: + """Resolve from explicit → env → default.""" + return explicit or os.environ.get(env_var, default) + + +@dataclass(frozen=True) +class ConfigEntry: + """A single resolvable configuration value.""" + + name: str + env_var: str + default: str + description: str = "" + + def resolve(self, explicit: str | None = None) -> str: + return resolve_config_value(explicit, self.env_var, self.default) diff --git a/alphaforge/data/public_web/cftc_cot.py b/alphaforge/data/public_web/cftc_cot.py index b67ff5a..b9f14e2 100644 --- a/alphaforge/data/public_web/cftc_cot.py +++ b/alphaforge/data/public_web/cftc_cot.py @@ -116,6 +116,24 @@ "13874A": "sp500_e_mini", "33874E": "sp500_micro", "209742": "vix", # VIX futures (alternative code) + # G10 FX futures (CME) + "099741": "eur", # Euro FX + "096742": "gbp", # British Pound + "097741": "jpy", # Japanese Yen + "092741": "chf", # Swiss Franc + "090741": "cad", # Canadian Dollar + "232741": "aud", # Australian Dollar + "112741": "nzd", # New Zealand Dollar + "095741": "mxn", # Mexican Peso (not G10 but heavily traded) + "089741": "sek", # Swedish Krona + "088741": "nok", # Norwegian Krone + # US rates futures + "13874P": "sofr_3m", # Three-Month SOFR (CME) + "134741": "ust_10y", # 10-Year T-Note + "020601": "ust_30y", # T-Bond (30Y) + "044601": "ust_5y", # 5-Year T-Note + "042601": "ust_2y", # 2-Year T-Note + "043602": "fed_funds", # 30-Day Federal Funds } # Regex fallbacks for market-name-based contract detection. @@ -123,6 +141,15 @@ (re.compile(r"\bVIX\b", re.IGNORECASE), "vix"), (re.compile(r"\bCBOE VOLATILITY INDEX\b", re.IGNORECASE), "vix"), (re.compile(r"\bS&P 500\b", re.IGNORECASE), "sp500"), + (re.compile(r"\bEURO FX\b", re.IGNORECASE), "eur"), + (re.compile(r"\bBRITISH POUND\b", re.IGNORECASE), "gbp"), + (re.compile(r"\bJAPANESE YEN\b", re.IGNORECASE), "jpy"), + (re.compile(r"\bSWISS FRANC\b", re.IGNORECASE), "chf"), + (re.compile(r"\bCANADIAN DOLLAR\b", re.IGNORECASE), "cad"), + (re.compile(r"\bAUSTRALIAN DOLLAR\b", re.IGNORECASE), "aud"), + (re.compile(r"\bNEW ZEALAND DOLLAR\b|\bNZ DOLLAR\b", re.IGNORECASE), "nzd"), + (re.compile(r"\bSOFR\b", re.IGNORECASE), "sofr_3m"), + (re.compile(r"\b10.YEAR\b.*\bT.NOTE\b", re.IGNORECASE), "ust_10y"), ] diff --git a/alphaforge/data/short_rates.py b/alphaforge/data/short_rates.py index e2e5e43..c5af127 100644 --- a/alphaforge/data/short_rates.py +++ b/alphaforge/data/short_rates.py @@ -3,7 +3,7 @@ from __future__ import annotations from dataclasses import dataclass, field -from typing import Mapping +from typing import Any, Mapping import pandas as pd @@ -105,7 +105,7 @@ def _fetch_fred_panel( ctx: DataContext, *, source: str, - series_map: Mapping[object, str], + series_map: Mapping[Any, str], start: pd.Timestamp, end: pd.Timestamp, sort_labels: bool = True, diff --git a/alphaforge/evaluation/__init__.py b/alphaforge/evaluation/__init__.py new file mode 100644 index 0000000..a607243 --- /dev/null +++ b/alphaforge/evaluation/__init__.py @@ -0,0 +1,48 @@ +"""Evaluation primitives — metric protocol and standard implementations. + +This package provides the foundational building blocks for forecast +evaluation. It is intentionally generic (not specific to nowcasting or +PIT data) so that downstream libraries like ``nowcast-data`` can compose +these primitives into domain-specific evaluation pipelines. + +Key components: + +- :class:`MetricFn` — a ``Protocol`` that any accuracy metric must satisfy. +- Five built-in metric classes: :class:`RMSE`, :class:`MAE`, + :class:`DirectionalAccuracy`, :class:`MAPE`, :class:`MeanError`. +- Two pre-built suites: :data:`DEFAULT_METRICS` (RMSE + MAE + DA) and + :data:`BENCHMARK_METRICS` (adds MeanError + MAPE). + +Downstream usage (nowcast-data):: + + from alphaforge.evaluation.metrics import BENCHMARK_METRICS + from nowcast_data.models.evaluation import benchmark_evaluation_suite + + results = benchmark_evaluation_suite( + predictions, + truth_definitions={"advance": ("y_true_release_1", 1)}, + metrics=list(BENCHMARK_METRICS), + ) +""" + +from .metrics import ( + BENCHMARK_METRICS, + DEFAULT_METRICS, + MAE, + MAPE, + RMSE, + DirectionalAccuracy, + MeanError, + MetricFn, +) + +__all__ = [ + "MetricFn", + "RMSE", + "MAE", + "DirectionalAccuracy", + "MAPE", + "MeanError", + "DEFAULT_METRICS", + "BENCHMARK_METRICS", +] diff --git a/alphaforge/evaluation/metrics.py b/alphaforge/evaluation/metrics.py new file mode 100644 index 0000000..43bdaa8 --- /dev/null +++ b/alphaforge/evaluation/metrics.py @@ -0,0 +1,311 @@ +"""Pluggable forecast accuracy metrics. + +This module defines a :class:`MetricFn` protocol and a library of standard +implementations that can be composed into any evaluation pipeline. The +protocol is ``runtime_checkable``, so ``isinstance(obj, MetricFn)`` works +for custom metrics without inheriting from a base class. + +Design rationale +~~~~~~~~~~~~~~~~ + +Published nowcasting benchmarks report different metric sets — GDPNow and +the Atlanta Fed use RMSE/MAE, the NY Fed reports Log Predictive Score, the +IMF reports Directional Accuracy, and the ECB uses CRPS. Rather than +hard-coding metric logic into evaluation functions, we define a thin protocol +and let callers compose their own metric suites. + +Creating a custom metric +~~~~~~~~~~~~~~~~~~~~~~~~ + +Any class satisfying the :class:`MetricFn` protocol works. The only +requirements are a ``name`` property (used as the column header in result +DataFrames) and a ``__call__(y_pred, y_true) -> float`` method:: + + class MedianAbsoluteError: + name = "median_ae" + + def __call__(self, y_pred, y_true): + return float(np.median(np.abs(y_pred - y_true))) + +Then pass it to any evaluation function:: + + from alphaforge.evaluation.metrics import RMSE + decompose_accuracy_by_horizon(predictions, metrics=[RMSE(), MedianAbsoluteError()]) + +Pre-built suites +~~~~~~~~~~~~~~~~ + +Two convenience tuples are provided: + +- :data:`DEFAULT_METRICS` — ``(RMSE, MAE, DirectionalAccuracy)`` for general + use where only basic accuracy is needed. +- :data:`BENCHMARK_METRICS` — adds ``MeanError`` (bias) and ``MAPE`` for + benchmark comparison tables that need to match published reporting. + +Examples +-------- +>>> import numpy as np +>>> from alphaforge.evaluation.metrics import RMSE, MAE, DirectionalAccuracy + +Compute a single metric: + +>>> rmse = RMSE() +>>> rmse(np.array([1.0, 2.0, 3.0]), np.array([1.0, 2.0, 4.0])) +0.5773502691896258 + +Check correct sign fraction: + +>>> da = DirectionalAccuracy() +>>> da(np.array([0.5, -0.3, 1.2]), np.array([0.1, 0.4, 0.8])) +0.6666666666666666 + +Use the protocol for runtime type checks: + +>>> from alphaforge.evaluation.metrics import MetricFn +>>> isinstance(RMSE(), MetricFn) +True + +See Also +-------- +nowcast_data.models.evaluation.decompose_accuracy_by_horizon : + Publication-date-anchored horizon decomposition that accepts ``MetricFn``. +nowcast_data.models.evaluation.benchmark_evaluation_suite : + Full metrics x truths x horizons evaluation. +nowcast_data.utils.metrics.compute_forecast_metrics : + Low-level metric computation that accepts optional ``MetricFn`` sequence. +""" + +from __future__ import annotations + +from typing import Protocol, runtime_checkable + +import numpy as np + +# --------------------------------------------------------------------------- +# Protocol +# --------------------------------------------------------------------------- + +@runtime_checkable +class MetricFn(Protocol): + """Protocol for forecast accuracy metrics. + + Any object with a ``name`` property and the correct call signature + satisfies this protocol. The ``@runtime_checkable`` decorator enables + ``isinstance(obj, MetricFn)`` checks at runtime, which evaluation + functions use to validate user-supplied metrics. + + Attributes + ---------- + name : str + Short, snake_case identifier used as the column header in result + DataFrames (e.g. ``"rmse"``, ``"directional_accuracy"``). + + Parameters (when called) + ------------------------ + y_pred : np.ndarray + 1-D array of predicted values. + y_true : np.ndarray + 1-D array of ground-truth values, same length as *y_pred*. + + Returns (when called) + --------------------- + float + Scalar metric value. Return ``np.nan`` when the metric is + undefined for the given inputs (e.g. MAPE with all-zero truths). + + Examples + -------- + Implement a custom metric: + + >>> class MedianAbsoluteError: + ... name = "median_ae" + ... def __call__(self, y_pred, y_true): + ... return float(np.median(np.abs(y_pred - y_true))) + >>> isinstance(MedianAbsoluteError(), MetricFn) + True + """ + + @property + def name(self) -> str: ... + + def __call__(self, y_pred: np.ndarray, y_true: np.ndarray) -> float: ... + + +# --------------------------------------------------------------------------- +# Built-in implementations +# --------------------------------------------------------------------------- + +class RMSE: + """Root Mean Squared Error. + + .. math:: + + \\text{RMSE} = \\sqrt{\\frac{1}{n} \\sum_{i=1}^{n} (\\hat{y}_i - y_i)^2} + + The standard accuracy metric in the nowcasting literature. Penalizes + large errors more than :class:`MAE` due to the squaring term. + + Examples + -------- + >>> RMSE()(np.array([1.0, 2.0, 3.0]), np.array([1.0, 2.0, 4.0])) + 0.5773502691896258 + """ + + name: str = "rmse" + + def __call__(self, y_pred: np.ndarray, y_true: np.ndarray) -> float: + return float(np.sqrt(np.mean((y_pred - y_true) ** 2))) + + def __repr__(self) -> str: + return "RMSE()" + + +class MAE: + """Mean Absolute Error. + + .. math:: + + \\text{MAE} = \\frac{1}{n} \\sum_{i=1}^{n} |\\hat{y}_i - y_i| + + More robust to outliers than :class:`RMSE`. GDPNow reports both + RMSE (1.17) and MAE (0.77) for 2011--2025. + + Examples + -------- + >>> MAE()(np.array([1.0, 2.0, 3.0]), np.array([1.0, 2.0, 4.0])) + 0.3333333333333333 + """ + + name: str = "mae" + + def __call__(self, y_pred: np.ndarray, y_true: np.ndarray) -> float: + return float(np.mean(np.abs(y_pred - y_true))) + + def __repr__(self) -> str: + return "MAE()" + + +class DirectionalAccuracy: + """Fraction of predictions with the correct sign. + + .. math:: + + \\text{DA} = \\frac{1}{n} \\sum_{i=1}^{n} + \\mathbb{1}[\\text{sign}(\\hat{y}_i) = \\text{sign}(y_i)] + + Critical for recession detection — tells you whether the model + correctly identifies positive vs. negative GDP growth. Reported by + the IMF (WP/2025/252) and ECB as a key evaluation criterion. + + Notes + ----- + When both ``y_pred`` and ``y_true`` are zero, ``np.sign`` returns 0 + for both, so the pair counts as a correct prediction. + + Examples + -------- + >>> DirectionalAccuracy()(np.array([1, -1, 1]), np.array([1, 1, -1])) + 0.3333333333333333 + """ + + name: str = "directional_accuracy" + + def __call__(self, y_pred: np.ndarray, y_true: np.ndarray) -> float: + correct = np.sign(y_pred) == np.sign(y_true) + return float(np.mean(correct)) + + def __repr__(self) -> str: + return "DirectionalAccuracy()" + + +class MAPE: + """Mean Absolute Percentage Error. + + .. math:: + + \\text{MAPE} = \\frac{1}{n} \\sum_{i=1}^{n} + \\left| \\frac{\\hat{y}_i - y_i}{y_i} \\right| + + Provides scale-independent accuracy, useful for cross-country + comparisons (e.g. IMF cross-country nowcast evaluations). + + Returns ``nan`` when all true values are near-zero (``|y_i| < 1e-10``), + since the metric is undefined in that case. + + Notes + ----- + Observations where ``|y_true| < 1e-10`` are excluded from the + computation to avoid division by zero. If *all* observations are + excluded, the result is ``nan``. + + Examples + -------- + >>> MAPE()(np.array([1.1, 2.2]), np.array([1.0, 2.0])) + 0.1 + >>> import math + >>> math.isnan(MAPE()(np.array([1.0]), np.array([0.0]))) + True + """ + + name: str = "mape" + + def __call__(self, y_pred: np.ndarray, y_true: np.ndarray) -> float: + mask = np.abs(y_true) > 1e-10 + if not mask.any(): + return float("nan") + return float(np.mean(np.abs((y_pred[mask] - y_true[mask]) / y_true[mask]))) + + def __repr__(self) -> str: + return "MAPE()" + + +class MeanError: + """Signed mean error (bias). + + .. math:: + + \\text{ME} = \\frac{1}{n} \\sum_{i=1}^{n} (\\hat{y}_i - y_i) + + Positive values indicate the forecast is systematically too high; + negative values indicate systematic under-prediction. An unbiased + forecast has ``MeanError ≈ 0``. + + Useful for diagnosing whether a model tends to over-predict or + under-predict GDP growth. + + Examples + -------- + >>> MeanError()(np.array([2.0, 3.0, 4.0]), np.array([1.0, 2.0, 3.0])) + 1.0 + """ + + name: str = "mean_error" + + def __call__(self, y_pred: np.ndarray, y_true: np.ndarray) -> float: + return float(np.mean(y_pred - y_true)) + + def __repr__(self) -> str: + return "MeanError()" + + +# --------------------------------------------------------------------------- +# Pre-built suites +# --------------------------------------------------------------------------- + +DEFAULT_METRICS: tuple[MetricFn, ...] = (RMSE(), MAE(), DirectionalAccuracy()) +"""Default metric suite: RMSE, MAE, and Directional Accuracy. + +Used by evaluation functions when no explicit ``metrics`` argument is +provided. Covers the two most common point-forecast accuracy measures +plus sign correctness. +""" + +BENCHMARK_METRICS: tuple[MetricFn, ...] = ( + RMSE(), MAE(), DirectionalAccuracy(), MeanError(), MAPE(), +) +"""Extended suite for benchmark comparison with published results. + +Adds :class:`MeanError` (bias detection) and :class:`MAPE` +(scale-independent accuracy) to the default suite. Matches the metric +set commonly reported across GDPNow, NY Fed, IMF, and ECB publications. +""" diff --git a/alphaforge/logging.py b/alphaforge/logging.py new file mode 100644 index 0000000..8357010 --- /dev/null +++ b/alphaforge/logging.py @@ -0,0 +1,107 @@ +"""Structured logging for alphaforge and downstream packages. + +Usage:: + + from alphaforge.logging import configure_logging, get_logger + + # Call once at startup + configure_logging(level="INFO", format="json") # or "console" + + # In any module + log = get_logger(__name__) + log.info("ingesting data", source="cftc_cot", rows=48) +""" +from __future__ import annotations + +import logging +import sys +from pathlib import Path + +import structlog + +# --------------------------------------------------------------------------- +# Standard context keys +# --------------------------------------------------------------------------- + +CTX_SOURCE = "source" +CTX_PIPELINE = "pipeline" +CTX_INSTRUMENT = "instrument" +CTX_ASOF = "asof" +CTX_TABLE = "table" +CTX_PHASE = "phase" +CTX_ROWS = "rows" +CTX_DURATION_MS = "duration_ms" + +# --------------------------------------------------------------------------- +# Public API +# --------------------------------------------------------------------------- + + +def configure_logging( + level: str = "INFO", + format: str = "console", + log_file: str | Path | None = None, +) -> None: + """Configure structured logging for the entire process. + + Parameters + ---------- + level : minimum log level + format : ``"console"`` for human-readable or ``"json"`` for JSON lines + log_file : optional path for JSON-line file output (append mode) + """ + shared_processors: list[structlog.types.Processor] = [ + structlog.contextvars.merge_contextvars, + structlog.stdlib.add_log_level, + structlog.stdlib.add_logger_name, + structlog.processors.TimeStamper(fmt="iso", utc=True), + structlog.processors.StackInfoRenderer(), + structlog.processors.format_exc_info, + ] + + if format == "json": + renderer: structlog.types.Processor = structlog.processors.JSONRenderer() + else: + renderer = structlog.dev.ConsoleRenderer() + + structlog.configure( + processors=[ + *shared_processors, + structlog.stdlib.ProcessorFormatter.wrap_for_formatter, + ], + wrapper_class=structlog.stdlib.BoundLogger, + context_class=dict, + logger_factory=structlog.stdlib.LoggerFactory(), + cache_logger_on_first_use=True, + ) + + formatter = structlog.stdlib.ProcessorFormatter( + processors=[ + structlog.stdlib.ProcessorFormatter.remove_processors_meta, + renderer, + ], + ) + + root = logging.getLogger() + root.handlers.clear() + root.setLevel(getattr(logging, level.upper(), logging.INFO)) + + stream_handler = logging.StreamHandler(sys.stderr) + stream_handler.setFormatter(formatter) + root.addHandler(stream_handler) + + if log_file is not None: + json_formatter = structlog.stdlib.ProcessorFormatter( + processors=[ + structlog.stdlib.ProcessorFormatter.remove_processors_meta, + structlog.processors.JSONRenderer(), + ], + ) + file_handler = logging.FileHandler(str(log_file), mode="a") + file_handler.setFormatter(json_formatter) + root.addHandler(file_handler) + + +def get_logger(name: str | None = None) -> structlog.stdlib.BoundLogger: + """Return a bound logger for *name*.""" + return structlog.get_logger(name) diff --git a/alphaforge/pipeline/__init__.py b/alphaforge/pipeline/__init__.py new file mode 100644 index 0000000..d0c7b61 --- /dev/null +++ b/alphaforge/pipeline/__init__.py @@ -0,0 +1,44 @@ +"""Composable pipeline component protocols and orchestrators.""" + +from .health import ( + CFTC_COT_HEALTH_POLICY, + DTCC_PPD_HEALTH_POLICY, + HealthPolicy, + HealthStatus, + SourceHealthPolicy, + SourceHealthStatus, + assess_health, + assess_source_health, +) +from .protocols import ( + Filter, + Parametric, + Pipeline, + PipelineVariant, + Signal, + SimplePipeline, + Transformer, +) +from .tracker import HealthTracker, SourceHealthTracker +from .weights import adjust_weights_for_health + +__all__ = [ + "Filter", + "Parametric", + "Pipeline", + "PipelineVariant", + "Signal", + "SimplePipeline", + "Transformer", + "SourceHealthPolicy", + "SourceHealthStatus", + "assess_source_health", + "CFTC_COT_HEALTH_POLICY", + "DTCC_PPD_HEALTH_POLICY", + "SourceHealthTracker", + "HealthPolicy", + "HealthStatus", + "HealthTracker", + "assess_health", + "adjust_weights_for_health", +] diff --git a/alphaforge/pipeline/health.py b/alphaforge/pipeline/health.py new file mode 100644 index 0000000..b4045b4 --- /dev/null +++ b/alphaforge/pipeline/health.py @@ -0,0 +1,144 @@ +"""Source health policy and assessment for data pipelines.""" +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING + +import pandas as pd + +if TYPE_CHECKING: + from alphaforge.pit.release_rules import ReleaseRule + + +@dataclass(frozen=True) +class SourceHealthPolicy: + """Configures how to respond when a data source goes stale.""" + + expected_cadence: pd.Timedelta + release_rule: ReleaseRule | None = None # precise schedule (optional) + grace_period: pd.Timedelta = pd.Timedelta(days=2) + stale_threshold: pd.Timedelta = pd.Timedelta(days=14) + dead_threshold: pd.Timedelta = pd.Timedelta(days=42) + weight_decay_start: pd.Timedelta | None = None # defaults to stale_threshold + weight_decay_half_life: pd.Timedelta = pd.Timedelta(days=7) + + +@dataclass(frozen=True) +class SourceHealthStatus: + """Health assessment for a single data source at a point in time.""" + + source_name: str + latest_obs_date: pd.Timestamp | None + asof: pd.Timestamp + age: pd.Timedelta | None + age_days: float | None + status: str # "ok", "late", "stale", "dead", "empty" + weight_factor: float + expected_next: pd.Timestamp | None + message: str + + +def assess_source_health( + source_name: str, + latest_obs_date: pd.Timestamp | None, + asof: pd.Timestamp, + policy: SourceHealthPolicy, +) -> SourceHealthStatus: + """Evaluate the health of a data source.""" + if latest_obs_date is None: + return SourceHealthStatus( + source_name=source_name, + latest_obs_date=None, + asof=asof, + age=None, + age_days=None, + status="empty", + weight_factor=0.0, + expected_next=None, + message=f"{source_name}: no observations found", + ) + + # Ensure both timestamps are comparable (both tz-aware or both naive) + _asof = asof + _latest = latest_obs_date + if _asof.tzinfo is not None and _latest.tzinfo is None: + _latest = _latest.tz_localize("UTC") + elif _asof.tzinfo is None and _latest.tzinfo is not None: + _asof = _asof.tz_localize("UTC") + + age = _asof - _latest + age_days = age.total_seconds() / 86400.0 + expected_next = _latest + policy.expected_cadence + + ok_limit = policy.expected_cadence + policy.grace_period + decay_start = ( + policy.weight_decay_start + if policy.weight_decay_start is not None + else policy.stale_threshold + ) + + if age <= ok_limit: + status = "ok" + weight = 1.0 + msg = f"{source_name}: OK (age {age_days:.0f}d)" + elif age <= policy.stale_threshold: + status = "late" + weight = 1.0 + msg = f"{source_name}: late (age {age_days:.0f}d, expected every {policy.expected_cadence.days}d)" + elif age <= policy.dead_threshold: + status = "stale" + decay_age = age - decay_start + hl = policy.weight_decay_half_life + weight = max(0.0, min(1.0, 2.0 ** (-decay_age / hl))) + msg = ( + f"{source_name}: stale (age {age_days:.0f}d, weight {weight:.2f})" + ) + else: + status = "dead" + weight = 0.0 + msg = ( + f"{source_name}: dead (age {age_days:.0f}d, " + f"exceeds dead threshold {policy.dead_threshold.days}d)" + ) + + return SourceHealthStatus( + source_name=source_name, + latest_obs_date=latest_obs_date, + asof=asof, + age=age, + age_days=age_days, + status=status, + weight_factor=weight, + expected_next=expected_next, + message=msg, + ) + + +# --------------------------------------------------------------------------- +# Default policies for known sources +# --------------------------------------------------------------------------- + +CFTC_COT_HEALTH_POLICY = SourceHealthPolicy( + expected_cadence=pd.Timedelta(days=7), + grace_period=pd.Timedelta(days=3), + stale_threshold=pd.Timedelta(days=21), + dead_threshold=pd.Timedelta(days=56), + weight_decay_half_life=pd.Timedelta(days=7), +) + +DTCC_PPD_HEALTH_POLICY = SourceHealthPolicy( + expected_cadence=pd.Timedelta(days=1), + grace_period=pd.Timedelta(days=3), + stale_threshold=pd.Timedelta(days=7), + dead_threshold=pd.Timedelta(days=30), + weight_decay_half_life=pd.Timedelta(days=3), +) + + +# --------------------------------------------------------------------------- +# Generalized aliases (Source prefix dropped for generic usage) +# --------------------------------------------------------------------------- + +HealthPolicy = SourceHealthPolicy +HealthStatus = SourceHealthStatus +assess_health = assess_source_health diff --git a/alphaforge/pipeline/protocols.py b/alphaforge/pipeline/protocols.py new file mode 100644 index 0000000..3eeb297 --- /dev/null +++ b/alphaforge/pipeline/protocols.py @@ -0,0 +1,138 @@ +"""Composable pipeline component protocols. + +Three stateless building blocks for research mode: + +- **Filter**: removes rows from a DataFrame. +- **Transformer**: adds or modifies columns. +- **Signal**: collapses a panel into a cross-sectional score. + +One marker for training mode: + +- **Parametric**: component exposes named parameters that can be + optimized. In research mode these are plain floats; in training + mode they can be replaced with tensor-backed values. + +Two concrete orchestrators: + +- **SimplePipeline**: linear chain of filters -> transformers -> signal. +- **PipelineVariant**: a named (filter config, pipeline) pair for model selection. +""" +from __future__ import annotations + +import time +from dataclasses import dataclass, field +from typing import Any, Protocol, runtime_checkable + +import pandas as pd + +from alphaforge.logging import get_logger + +_log = get_logger(__name__) + + +@runtime_checkable +class Filter(Protocol): + """Remove rows from a DataFrame based on criteria.""" + + name: str + + def apply(self, df: pd.DataFrame, **params: Any) -> pd.DataFrame: ... + + +@runtime_checkable +class Transformer(Protocol): + """Add or modify columns in a DataFrame.""" + + name: str + + def transform(self, df: pd.DataFrame, **params: Any) -> pd.DataFrame: ... + + +@runtime_checkable +class Signal(Protocol): + """Collapse a panel into a cross-sectional score (entity -> float).""" + + name: str + + def score(self, df: pd.DataFrame, **params: Any) -> pd.Series: ... + + +@runtime_checkable +class Parametric(Protocol): + """Marker for components with tunable parameters. + + Transformers and Signals may implement this. Filters do NOT. + """ + + name: str + + def get_params(self) -> dict[str, float]: ... + def set_params(self, params: dict[str, float]) -> None: ... + + +@runtime_checkable +class Pipeline(Protocol): + """Ordered sequence of components with a terminal signal.""" + + name: str + + def run(self, df: pd.DataFrame, **params: Any) -> pd.Series: ... + + def get_all_params(self) -> dict[str, dict[str, float]]: ... + + def set_all_params(self, params: dict[str, dict[str, float]]) -> None: ... + + +@dataclass +class SimplePipeline: + """Linear chain: filters -> transformers -> signal.""" + + name: str + filters: list[Filter] = field(default_factory=list) + transformers: list[Transformer] = field(default_factory=list) + signal: Signal | None = None + + def run(self, df: pd.DataFrame, **params: Any) -> pd.Series: + t0 = time.monotonic() + out = df + for f in self.filters: + before = len(out) + out = f.apply(out, **params) + _log.debug("filter_applied", pipeline=self.name, filter=f.name, before=before, after=len(out)) + if out.empty: + _log.info("pipeline_empty_after_filter", pipeline=self.name, filter=f.name) + return pd.Series(dtype="float64") + for t in self.transformers: + out = t.transform(out, **params) + if self.signal is None: + raise ValueError(f"Pipeline {self.name!r} has no terminal Signal") + result = self.signal.score(out, **params) + duration_ms = int((time.monotonic() - t0) * 1000) + _log.debug("pipeline_run_complete", pipeline=self.name, entities=len(result), duration_ms=duration_ms) + return result + + def get_all_params(self) -> dict[str, dict[str, float]]: + result: dict[str, dict[str, float]] = {} + for component in [*self.filters, *self.transformers]: + if isinstance(component, Parametric): + result[component.name] = component.get_params() + if self.signal is not None and isinstance(self.signal, Parametric): + result[self.signal.name] = self.signal.get_params() + return result + + def set_all_params(self, params: dict[str, dict[str, float]]) -> None: + for component in [*self.filters, *self.transformers]: + if isinstance(component, Parametric) and component.name in params: + component.set_params(params[component.name]) + if self.signal is not None and isinstance(self.signal, Parametric): + if self.signal.name in params: + self.signal.set_params(params[self.signal.name]) + + +@dataclass(frozen=True) +class PipelineVariant: + """A named (filter config, pipeline) pair for model selection.""" + + name: str + pipeline: SimplePipeline + description: str = "" diff --git a/alphaforge/pipeline/torch_protocols.py b/alphaforge/pipeline/torch_protocols.py new file mode 100644 index 0000000..2f76514 --- /dev/null +++ b/alphaforge/pipeline/torch_protocols.py @@ -0,0 +1,148 @@ +"""Torch-compatible pipeline component protocols. + +These extend the base Parametric protocol for components that can +participate in gradient-based optimization. +""" +from __future__ import annotations + +from typing import Protocol, runtime_checkable + +import numpy as np + +try: + import torch + import torch.nn as nn + + HAS_TORCH = True +except ImportError: # pragma: no cover + HAS_TORCH = False + + +@runtime_checkable +class TorchParametric(Protocol): + """A component whose parameters are torch tensors.""" + + def parameters(self) -> list: ... # list[nn.Parameter] + def named_parameters(self) -> list: ... # list[tuple[str, nn.Parameter]] + def get_params(self) -> dict[str, float]: ... + def set_params(self, params: dict[str, float]) -> None: ... + + +if HAS_TORCH: + + class DifferentiableSignalModule(nn.Module): + """Base class for differentiable signal pipeline components.""" + + name: str = "base" + + def forward(self, x: torch.Tensor) -> torch.Tensor: + raise NotImplementedError + + def get_params(self) -> dict[str, float]: + result = {} + for name, p in self.named_parameters(): + if not p.requires_grad: + continue + if p.numel() == 1: + result[name] = float(p.item()) + else: + # Multi-element params: store as first element (summary) + result[name] = float(p[0].item()) + return result + + def set_params(self, params: dict[str, float]) -> None: + with torch.no_grad(): + for name, value in params.items(): + parts = name.split(".") + obj = self + for part in parts[:-1]: + obj = getattr(obj, part) + if hasattr(obj, parts[-1]): + param = getattr(obj, parts[-1]) + if isinstance(param, nn.Parameter): + param.fill_(value) + + class _NumpyAutograd(torch.autograd.Function): + """Custom autograd bridging numpy -> torch.""" + + @staticmethod + def forward(ctx, x, forward_fn, backward_fn, params_dict, fd_eps): + x_np = x.detach().cpu().numpy() + params_float = {k: float(v.item()) for k, v in params_dict.items()} + result_np = forward_fn(x_np, params_float) + ctx.save_for_backward(x) + ctx.forward_fn = forward_fn + ctx.backward_fn = backward_fn + ctx.params_dict = params_dict + ctx.params_float = params_float + ctx.fd_eps = fd_eps + return torch.tensor(result_np, dtype=x.dtype, device=x.device) + + @staticmethod + def backward(ctx, grad_output): + (x,) = ctx.saved_tensors + x_np = x.detach().cpu().numpy() + grad_np = grad_output.detach().cpu().numpy() + + if ctx.backward_fn is not None: + grads = ctx.backward_fn(x_np, grad_np, ctx.params_float) + grad_x = torch.tensor( + grads.get("input", np.zeros_like(x_np)), + dtype=x.dtype, + device=x.device, + ) + else: + # Finite differences for input gradient + eps = ctx.fd_eps + grad_x = torch.zeros_like(x) + flat = x_np.ravel() + for i in range(min(len(flat), 100)): # cap for perf + flat_p = flat.copy() + flat_p[i] += eps + x_p = flat_p.reshape(x_np.shape) + out_p = ctx.forward_fn(x_p, ctx.params_float) + flat_m = flat.copy() + flat_m[i] -= eps + x_m = flat_m.reshape(x_np.shape) + out_m = ctx.forward_fn(x_m, ctx.params_float) + deriv = (out_p - out_m) / (2 * eps) + grad_x.view(-1)[i] = ( + torch.tensor(deriv, dtype=x.dtype).ravel() * grad_output.ravel() + ).sum() + + return grad_x, None, None, None, None + + class NumpyDifferentiableWrapper(DifferentiableSignalModule): + """Wrap a numpy function as a differentiable torch module.""" + + def __init__( + self, + forward_fn, + backward_fn=None, + param_names: list[str] | None = None, + init_params: dict[str, float] | None = None, + fd_eps: float = 1e-5, + name: str = "numpy_wrapper", + ): + super().__init__() + self.name = name + self.forward_fn = forward_fn + self.backward_fn = backward_fn + self.fd_eps = fd_eps + self.param_names_list = list(param_names or []) + + for pname in self.param_names_list: + val = (init_params or {}).get(pname, 0.0) + setattr(self, pname, nn.Parameter(torch.tensor(float(val)))) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + params = { + pname: getattr(self, pname) for pname in self.param_names_list + } + return _NumpyAutograd.apply( + x, + self.forward_fn, + self.backward_fn, + params, + self.fd_eps, + ) diff --git a/alphaforge/pipeline/tracker.py b/alphaforge/pipeline/tracker.py new file mode 100644 index 0000000..7efd0d1 --- /dev/null +++ b/alphaforge/pipeline/tracker.py @@ -0,0 +1,135 @@ +"""Source health tracker — persists health assessments in PIT.""" +from __future__ import annotations + +from dataclasses import dataclass + +import pandas as pd + +from alphaforge.pit.accessor import PITAccessor + +from .health import SourceHealthPolicy, SourceHealthStatus, assess_source_health + +_STATUS_CODE = {"ok": 0, "late": 1, "stale": 2, "dead": 3, "empty": 4} + + +@dataclass +class SourceHealthTracker: + """Track source health over time in the PIT store.""" + + pit: PITAccessor + policies: dict[str, SourceHealthPolicy] + + # Mapping from source_name to series_key prefix used to detect latest obs + source_prefixes: dict[str, str] | None = None + + def _prefix(self, source_name: str) -> str: + if self.source_prefixes and source_name in self.source_prefixes: + return self.source_prefixes[source_name] + return source_name.replace("_", ".") + + def _latest_obs_date(self, source_name: str) -> pd.Timestamp | None: + try: + result = self.pit.conn.execute( + "SELECT MAX(obs_date) FROM pit_observations WHERE source = ?", + [source_name], + ).fetchone() + except Exception: + return None + if result is None or result[0] is None: + return None + ts = pd.Timestamp(result[0]) + if ts.tzinfo is None: + ts = ts.tz_localize("UTC") + return ts + + def assess(self, source_name: str, asof: pd.Timestamp) -> SourceHealthStatus: + """Assess current health of a source by querying PIT for latest obs.""" + if source_name not in self.policies: + return SourceHealthStatus( + source_name=source_name, + latest_obs_date=None, + asof=asof, + age=None, + age_days=None, + status="empty", + weight_factor=0.0, + expected_next=None, + message=f"{source_name}: no policy configured", + ) + latest = self._latest_obs_date(source_name) + return assess_source_health(source_name, latest, asof, self.policies[source_name]) + + def assess_all(self, asof: pd.Timestamp) -> dict[str, SourceHealthStatus]: + """Assess all registered sources.""" + return {name: self.assess(name, asof) for name in self.policies} + + def record(self, status: SourceHealthStatus) -> None: + """Persist a health assessment into PIT.""" + base = f"health.{status.source_name}" + asof = status.asof + obs_date = asof + + rows = [ + { + "series_key": f"{base}.status_code", + "obs_date": obs_date, + "asof_utc": asof, + "value": float(_STATUS_CODE.get(status.status, 4)), + "source": "health", + }, + { + "series_key": f"{base}.weight_factor", + "obs_date": obs_date, + "asof_utc": asof, + "value": status.weight_factor, + "source": "health", + }, + ] + if status.age_days is not None: + rows.append({ + "series_key": f"{base}.age_days", + "obs_date": obs_date, + "asof_utc": asof, + "value": status.age_days, + "source": "health", + }) + + df = pd.DataFrame(rows) + self.pit.upsert_pit_observations(df, strict="coerce") + + def history( + self, + source_name: str, + start: pd.Timestamp | None = None, + end: pd.Timestamp | None = None, + ) -> pd.DataFrame: + """Retrieve health assessment history for a source.""" + prefix = f"health.{source_name}.%" + params: list = [prefix] + filters = ["series_key LIKE ?"] + + if start is not None: + filters.append("obs_date >= ?") + params.append(start) + if end is not None: + filters.append("obs_date <= ?") + params.append(end) + + where = " AND ".join(filters) + result = self.pit.conn.execute( + f"SELECT series_key, obs_date, asof_utc, value FROM pit_observations WHERE {where} ORDER BY obs_date", + params, + ).fetchdf() + return result + + def is_any_degraded(self, asof: pd.Timestamp) -> bool: + """Quick check: is any source late/stale/dead?""" + for name in self.policies: + status = self.assess(name, asof) + if status.status in ("late", "stale", "dead", "empty"): + return True + return False + + +# Generalized alias +HealthTracker = SourceHealthTracker diff --git a/alphaforge/pipeline/weights.py b/alphaforge/pipeline/weights.py new file mode 100644 index 0000000..abfe87a --- /dev/null +++ b/alphaforge/pipeline/weights.py @@ -0,0 +1,42 @@ +"""Staleness-aware weight adjustment for multi-source signal combination.""" +from __future__ import annotations + +from typing import Callable + + +def adjust_weights_for_health( + weights: dict[str, float], + health_fn: Callable[[str], float], + source_map: dict[str, str], + renormalize: bool = True, +) -> dict[str, float]: + """Adjust pipeline weights based on source health weight factors. + + Parameters + ---------- + weights : Nominal weights per pipeline name. + health_fn : Callable that takes a source name and returns a weight factor + in [0, 1]. Typically ``ctx.source_weight``. + source_map : Mapping from pipeline name to source name. + renormalize : If True, scale surviving weights so total matches original. + + Returns + ------- + Effective weights after health adjustment (and optional renormalization). + """ + effective: dict[str, float] = dict(weights) + + for name in list(effective): + source = source_map.get(name) + if source is not None: + factor = health_fn(source) + effective[name] = effective[name] * factor + + if renormalize: + original_total = sum(weights.values()) + current_total = sum(effective.values()) + if current_total > 0 and current_total < original_total: + scale = original_total / current_total + effective = {k: v * scale for k, v in effective.items()} + + return effective diff --git a/alphaforge/pit/__init__.py b/alphaforge/pit/__init__.py index 7ab2177..4e0fdec 100644 --- a/alphaforge/pit/__init__.py +++ b/alphaforge/pit/__init__.py @@ -42,6 +42,12 @@ coerce_pipeline_spec, ) from .ref_entity import make_ref_entity_id, parse_ref_entity_id +from .resolvers import ( + FrozenResolver, + LatestResolver, + RealtimeResolver, + VintageResolver, +) from .tasks import ( build_snapshot_tape, first_vintage_snapshot, @@ -60,6 +66,7 @@ ) from .transforms import PITTransformResult, PITTransformSpec from .validation import PITValidationReport, validate_pit_observations +from .views import VintageView __all__ = [ "PITAccessor", @@ -122,4 +129,9 @@ "build_snapshot_tape", "make_ref_entity_id", "parse_ref_entity_id", + "VintageView", + "VintageResolver", + "RealtimeResolver", + "LatestResolver", + "FrozenResolver", ] diff --git a/alphaforge/pit/missingness.py b/alphaforge/pit/missingness.py new file mode 100644 index 0000000..235912a --- /dev/null +++ b/alphaforge/pit/missingness.py @@ -0,0 +1,152 @@ +"""Missingness taxonomy and classification for nowcasting panels. + +Provides a shared classifier so Phase 3 (views) and Phase 4 (imputation) +use identical logic to determine *why* a cell is NaN. +""" + +from __future__ import annotations + +from datetime import date +from enum import Enum +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from alphaforge.pit.release_rules import ReleaseRule + + +class MissingnessReason(str, Enum): + """Why a panel cell is NaN.""" + + STRUCTURAL = "structural" + """Frequency mismatch — e.g. quarterly obs_date in a monthly panel.""" + + FUTURE = "future" + """The observation period has not yet ended (obs_date > asof_date).""" + + RAGGED_EDGE = "ragged_edge" + """Not yet published — the expected release date is after asof_date.""" + + TRUE_MISSING = "true_missing" + """Should have been available by now but is absent — data issue.""" + + +def classify_missingness( + *, + obs_date: date, + asof_date: date, + series_frequency: str, + panel_frequency: str = "M", + release_rule: ReleaseRule | None = None, + publication_lag_months: int | None = None, + realized_release_date: date | None = None, +) -> MissingnessReason: + """Classify why a panel cell is NaN. + + Evaluation follows a strict order: + + 1. **STRUCTURAL** — frequency mismatch (e.g. quarterly series in a + monthly panel has NaN for non-quarter-end months). + 2. **FUTURE** — ``obs_date`` is after ``asof_date``. + 3. **RAGGED_EDGE vs TRUE_MISSING** — compare ``asof_date`` against the + expected (or realized) publication date: + + * If a ``realized_release_date`` is provided (from AlphaForge PIT + storage), it takes precedence over the expected schedule. + * Otherwise, ``release_rule.expected_release_date(obs_date)`` is used. + * As a last fallback, ``publication_lag_months`` gives a coarse + estimate. + * If none are available, default to ``TRUE_MISSING``. + + Parameters + ---------- + obs_date + The observation / reference-period end date of the NaN cell. + asof_date + The "as-of" date of the vintage snapshot. + series_frequency + Frequency code of the series (``"Q"``, ``"M"``, ``"W"``, etc.). + panel_frequency + Frequency code of the panel grid (default ``"M"`` for monthly). + release_rule + Optional structured release schedule from catalog metadata. + publication_lag_months + Optional simple lag heuristic (backward compat). + realized_release_date + Optional realized publication date from AlphaForge. When provided + this always takes precedence over the expected schedule. + """ + # 1. Structural: quarterly series in a monthly panel at non-quarter-end + if _is_structural(obs_date, series_frequency, panel_frequency): + return MissingnessReason.STRUCTURAL + + # 2. Future: observation period hasn't ended + if obs_date > asof_date: + return MissingnessReason.FUTURE + + # 3. Ragged edge vs true missing + expected = _resolve_expected_date( + obs_date, + release_rule=release_rule, + publication_lag_months=publication_lag_months, + realized_release_date=realized_release_date, + ) + + if expected is None: + # No release information at all — assume it should be here + return MissingnessReason.TRUE_MISSING + + if asof_date < expected: + return MissingnessReason.RAGGED_EDGE + + return MissingnessReason.TRUE_MISSING + + +# --------------------------------------------------------------------------- +# Internal helpers +# --------------------------------------------------------------------------- + +_QUARTER_END_MONTHS = {3, 6, 9, 12} + + +def _is_structural( + obs_date: date, series_frequency: str, panel_frequency: str +) -> bool: + """Return True when the NaN is a frequency-mismatch artifact.""" + sf = series_frequency.upper() + pf = panel_frequency.upper() + + if sf == "Q" and pf == "M": + # Quarterly series only has values at quarter-end months + return obs_date.month not in _QUARTER_END_MONTHS + + return False + + +def _resolve_expected_date( + obs_date: date, + *, + release_rule: ReleaseRule | None, + publication_lag_months: int | None, + realized_release_date: date | None, +) -> date | None: + """Determine the best-available expected publication date. + + Precedence: + 1. realized_release_date (from AlphaForge PIT storage) + 2. release_rule.expected_release_date(obs_date) + 3. obs_date + publication_lag_months (coarse fallback) + 4. None (no information) + """ + if realized_release_date is not None: + return realized_release_date + + if release_rule is not None: + return release_rule.expected_release_date(obs_date) + + if publication_lag_months is not None: + m = obs_date.month + publication_lag_months + y = obs_date.year + (m - 1) // 12 + m = (m - 1) % 12 + 1 + return date(y, m, 1) + + return None diff --git a/alphaforge/pit/panel.py b/alphaforge/pit/panel.py new file mode 100644 index 0000000..4e40b2d --- /dev/null +++ b/alphaforge/pit/panel.py @@ -0,0 +1,111 @@ +"""PIT panel building utilities. + +Generic functions for assembling aligned panels from PIT snapshots. +Used by positioning (SignalContext) and nowcast-data (benchmark PIT builder). +""" +from __future__ import annotations + +from typing import Sequence + +import pandas as pd + +from alphaforge.pit.accessor import PITAccessor + + +def build_pit_panel( + pit: PITAccessor, + series_keys: dict[str, str], + asof: pd.Timestamp, + start: pd.Timestamp | None = None, + end: pd.Timestamp | None = None, + align_freq: str | None = None, +) -> pd.DataFrame: + """Build an aligned wide panel from PIT snapshots. + + Parameters + ---------- + pit : PITAccessor + series_keys : mapping from desired column name to PIT series_key + asof : point-in-time cutoff + start, end : date range + align_freq : if set, reindex to this frequency (with ffill for slower series) + + Returns + ------- + Wide DataFrame with index=obs_date, columns=column_name. + """ + series: dict[str, pd.Series] = {} + for col_name, skey in series_keys.items(): + snap = pit.get_snapshot(skey, asof, start=start, end=end) + series[col_name] = snap + + if not series: + return pd.DataFrame() + + df = pd.DataFrame(series) + + if align_freq is not None: + new_idx = pd.date_range( + start=df.index.min(), end=df.index.max(), freq=align_freq + ) + df = df.reindex(new_idx).ffill() + + df.index.name = "obs_date" + return df + + +def build_pit_panel_long( + pit: PITAccessor, + series_specs: Sequence[dict], + asof: pd.Timestamp, + start: pd.Timestamp | None = None, + end: pd.Timestamp | None = None, +) -> pd.DataFrame: + """Build a long-format panel (obs_date, entity_id, value columns). + + Parameters + ---------- + series_specs : list of dicts with keys: + - name: column alias + - series_key: PIT series key + - entity_id: (optional) entity identifier + asof, start, end : PIT query params + + Returns + ------- + Long DataFrame with columns: obs_date, entity_id, and one column per spec name. + """ + rows: list[dict] = [] + + for spec in series_specs: + name = spec["name"] + skey = spec["series_key"] + entity_id = spec.get("entity_id", name) + + snap = pit.get_snapshot(skey, asof, start=start, end=end) + for d, v in snap.items(): + rows.append({ + "obs_date": d, + "entity_id": entity_id, + name: v, + }) + + if not rows: + return pd.DataFrame(columns=["obs_date", "entity_id"]) + + return pd.DataFrame(rows) + + +def long_to_wide( + df: pd.DataFrame, + index_col: str = "obs_date", + columns_col: str = "series_key", + values_col: str = "value", +) -> pd.DataFrame: + """Pivot long PIT DataFrame to wide format.""" + return df.pivot_table( + index=index_col, + columns=columns_col, + values=values_col, + aggfunc="first", + ) diff --git a/alphaforge/pit/release_rules.py b/alphaforge/pit/release_rules.py new file mode 100644 index 0000000..ff8e30b --- /dev/null +++ b/alphaforge/pit/release_rules.py @@ -0,0 +1,280 @@ +"""Release schedule rules for macro series publication timing. + +Each rule class models a real-world release schedule pattern and exposes +``expected_release_date(obs_date)`` to compute the day a value is expected +to become publicly available. + +These rules are an *expectation layer* — realized PIT timestamps from +AlphaForge always take precedence when available. +""" + +from __future__ import annotations + +from abc import ABC, abstractmethod +from dataclasses import dataclass +from datetime import date +from typing import Any + +import pandas as pd +from pandas.tseries.holiday import USFederalHolidayCalendar +from pandas.tseries.offsets import CustomBusinessDay + +_US_BD = CustomBusinessDay(calendar=USFederalHolidayCalendar()) + +# --------------------------------------------------------------------------- +# Registry for tagged-union YAML parsing (populated by @_register below) +# --------------------------------------------------------------------------- + +RULE_REGISTRY: dict[str, type] = {} + + +# --------------------------------------------------------------------------- +# Base +# --------------------------------------------------------------------------- + + +@dataclass(frozen=True) +class ReleaseRule(ABC): + """Base class for publication schedule rules.""" + + rule_type: str = "" # overridden by subclass class-var + + @abstractmethod + def expected_release_date( + self, obs_date: date, release_number: int | None = None + ) -> date: + """Return the expected publication date for a given observation date. + + Parameters + ---------- + obs_date : date + The observation / reference-period end date. + release_number : int | None + For multi-release series (e.g. GDP), selects which release + (1 = advance, 2 = preliminary, …). Ignored by single-release + rules. + """ + + # Serialization helpers for round-trip YAML + def to_dict(self) -> dict[str, Any]: + """Serialize to a dict suitable for YAML output.""" + d: dict[str, Any] = {"type": self.rule_type} + for k, v in self.__dict__.items(): + if k == "rule_type": + continue + d[k] = v + return d + + @staticmethod + def from_dict(d: dict[str, Any]) -> ReleaseRule: + """Reconstruct a ReleaseRule from a YAML-style dict.""" + kwargs = dict(d) + type_key = kwargs.pop("type") + cls = RULE_REGISTRY.get(type_key) + if cls is None: + raise ValueError( + f"Unknown release rule type '{type_key}'. " + f"Known types: {sorted(RULE_REGISTRY)}" + ) + return cls(**kwargs) + + +def _register(cls: type[ReleaseRule]) -> type[ReleaseRule]: + """Decorator that adds a ReleaseRule subclass to the registry.""" + RULE_REGISTRY[cls.rule_type] = cls + return cls + + +# --------------------------------------------------------------------------- +# Concrete rule classes +# --------------------------------------------------------------------------- + + +def _month_start(ref: date, offset_months: int) -> date: + """Return the 1st day of the month that is *offset_months* after *ref*.""" + m = ref.month + offset_months + y = ref.year + (m - 1) // 12 + m = (m - 1) % 12 + 1 + return date(y, m, 1) + + +def _nth_business_day(anchor: date, n: int) -> date: + """Return the *n*-th US business day on or after *anchor*.""" + ts = pd.Timestamp(anchor) + # Move to the n-th business day (1-indexed) + result = ts + (n - 1) * _US_BD + # If anchor itself is not a business day, the offset already handles it + # but we need to make sure we start counting from the first valid day + first_bd = ts + 0 * _US_BD # snap to first business day on or after + result = first_bd + (n - 1) * _US_BD + return result.date() + + +@_register +@dataclass(frozen=True) +class NthBusinessDay(ReleaseRule): + """Published on the n-th US business day relative to an anchor month. + + Example: Employment Situation — 1st business day of the following month. + """ + + rule_type: str = "nth_business_day" + n: int = 1 + anchor: str = "following_month" # "following_month" | "same_month" + + def expected_release_date( + self, obs_date: date, release_number: int | None = None + ) -> date: + if self.anchor == "following_month": + start = _month_start(obs_date, 1) + elif self.anchor == "same_month": + start = _month_start(obs_date, 0) + else: + raise ValueError(f"Unknown anchor: {self.anchor}") + return _nth_business_day(start, self.n) + + +@_register +@dataclass(frozen=True) +class NthWeekday(ReleaseRule): + """Published on the n-th occurrence of a weekday in the anchor month. + + Example: CPI — typically released around the 10th–15th business day + of the following month, often modeled as ~2nd or 3rd week. + """ + + rule_type: str = "nth_weekday" + n: int = 1 + weekday: str = "Friday" # Monday–Sunday + anchor: str = "following_month" + + def expected_release_date( + self, obs_date: date, release_number: int | None = None + ) -> date: + _weekday_map = { + "Monday": 0, + "Tuesday": 1, + "Wednesday": 2, + "Thursday": 3, + "Friday": 4, + "Saturday": 5, + "Sunday": 6, + } + target_wd = _weekday_map[self.weekday] + + if self.anchor == "following_month": + first = _month_start(obs_date, 1) + elif self.anchor == "same_month": + first = _month_start(obs_date, 0) + else: + raise ValueError(f"Unknown anchor: {self.anchor}") + + # Find the first occurrence of target weekday in the month + days_ahead = (target_wd - first.weekday()) % 7 + first_occurrence = first + pd.Timedelta(days=days_ahead) + # Advance to the n-th occurrence + result = first_occurrence + pd.Timedelta(weeks=self.n - 1) + return result.date() if isinstance(result, pd.Timestamp) else result + + +@_register +@dataclass(frozen=True) +class CalendarDay(ReleaseRule): + """Published on a specific calendar day of the anchor month. + + Example: "15th of the following month." + """ + + rule_type: str = "calendar_day" + day: int = 15 + anchor: str = "following_month" + + def expected_release_date( + self, obs_date: date, release_number: int | None = None + ) -> date: + if self.anchor == "following_month": + start = _month_start(obs_date, 1) + elif self.anchor == "same_month": + start = _month_start(obs_date, 0) + elif self.anchor == "two_months_later": + start = _month_start(obs_date, 2) + else: + raise ValueError(f"Unknown anchor: {self.anchor}") + return start.replace(day=self.day) + + +@_register +@dataclass(frozen=True) +class FixedLagMonths(ReleaseRule): + """Published a fixed number of months after the observation month. + + Fallback rule for series where the exact schedule is not worth encoding. + """ + + rule_type: str = "fixed_lag_months" + months: int = 1 + + def expected_release_date( + self, obs_date: date, release_number: int | None = None + ) -> date: + return _month_start(obs_date, self.months) + + +@_register +@dataclass(frozen=True) +class QuarterlyRelease(ReleaseRule): + """GDP-style multi-release schedule. + + Models advance / preliminary / final releases with configurable lag + (in months after the quarter-end). + """ + + rule_type: str = "quarterly_release" + advance_lag_months: int = 1 + preliminary_lag_months: int = 2 + final_lag_months: int = 3 + + def expected_release_date( + self, obs_date: date, release_number: int | None = None + ) -> date: + rn = release_number or 1 + if rn == 1: + lag = self.advance_lag_months + elif rn == 2: + lag = self.preliminary_lag_months + else: + lag = self.final_lag_months + return _month_start(obs_date, lag) + + +@_register +@dataclass(frozen=True) +class WeeklyRelease(ReleaseRule): + """Weekly series released on a fixed weekday with a lag. + + Example: Initial Claims — released Thursday for the prior Saturday week. + """ + + rule_type: str = "weekly" + release_weekday: str = "Thursday" + lag_days: int = 5 + + def expected_release_date( + self, obs_date: date, release_number: int | None = None + ) -> date: + return obs_date + pd.Timedelta(days=self.lag_days) + + +@_register +@dataclass(frozen=True) +class CustomRule(ReleaseRule): + """Free-text description for exotic or not-yet-modeled schedules.""" + + rule_type: str = "custom" + description: str = "" + approximate_lag_months: int = 1 + + def expected_release_date( + self, obs_date: date, release_number: int | None = None + ) -> date: + return _month_start(obs_date, self.approximate_lag_months) diff --git a/alphaforge/pit/resolvers.py b/alphaforge/pit/resolvers.py new file mode 100644 index 0000000..034c027 --- /dev/null +++ b/alphaforge/pit/resolvers.py @@ -0,0 +1,156 @@ +"""Vintage resolvers for point-in-time data projection. + +A :class:`VintageResolver` translates a :class:`~alphaforge.pit.views.VintageView` +declaration into concrete ``asof_date`` resolution logic. The resolver sits +between the caller (e.g. a backtest loop) and the PIT adapter: the caller asks +"what should I fetch for this (series, obs_date, asof_date) triple?", and the +resolver returns the *effective* asof_date to pass to the adapter. +""" + +from __future__ import annotations + +from datetime import date +from typing import Protocol, runtime_checkable + +from alphaforge.pit.views import VintageView + +# --------------------------------------------------------------------------- +# Protocol +# --------------------------------------------------------------------------- + + +@runtime_checkable +class VintageResolver(Protocol): + """Resolves the effective asof_date for a fetch given a vintage view. + + Implementations are stateless (:class:`RealtimeResolver`, + :class:`LatestResolver`) or carry an immutable pre-computed revision + map (:class:`FrozenResolver`). + """ + + @property + def view(self) -> VintageView: + """The view this resolver implements.""" + ... + + def resolve( + self, + series_key: str, + obs_date: date, + requested_asof: date, + has_pit: bool, + ) -> date: + """Return the effective asof_date the adapter should use. + + Args: + series_key: Canonical series identifier. + obs_date: Observation date being resolved. + requested_asof: The walk-forward asof_date from the backtest. + has_pit: Whether this series has vintage history. + + Returns: + The asof_date to pass to ``adapter.fetch_asof()``. + """ + ... + + +# --------------------------------------------------------------------------- +# Concrete resolvers +# --------------------------------------------------------------------------- + + +class RealtimeResolver: + """Identity projection — returns ``requested_asof`` unchanged. + + Stateless. Zero overhead. + """ + + def __init__(self) -> None: + self._view = VintageView.realtime() + + @property + def view(self) -> VintageView: + return self._view + + def resolve( + self, + series_key: str, + obs_date: date, + requested_asof: date, + has_pit: bool, + ) -> date: + return requested_asof + + +class LatestResolver: + """Collapses all vintages to the most recent available. + + Stateless. Uses a far-future sentinel to exploit the adapter's + existing "latest vintage ≤ asof" semantics. + """ + + _SENTINEL = date(2099, 12, 31) + + def __init__(self) -> None: + self._view = VintageView.latest() + + @property + def view(self) -> VintageView: + return self._view + + def resolve( + self, + series_key: str, + obs_date: date, + requested_asof: date, + has_pit: bool, + ) -> date: + return self._SENTINEL if has_pit else requested_asof + + +class FrozenResolver: + """Resolves to the *n*-th release vintage for each (series, obs_date). + + Stateful: carries an immutable pre-computed revision map built at + construction time. After construction, :meth:`resolve` is a dict + lookup with no adapter calls. + + The revision map is built *outside* this class (typically by the + backtest harness which has access to the PIT adapter) and passed in + as a plain ``dict``. This keeps the resolver free of adapter + dependencies. + + Args: + revision_map: Mapping from ``(series_key, obs_date)`` to a + **sorted** list of vintage dates on which that observation + was revised. + n_releases: Which release to freeze to (1-indexed). Defaults to 3. + """ + + def __init__( + self, + revision_map: dict[tuple[str, date], list[date]], + n_releases: int = 3, + ) -> None: + self._view = VintageView.frozen(n=n_releases) + self._n = n_releases + self._revision_map = revision_map + + @property + def view(self) -> VintageView: + return self._view + + def resolve( + self, + series_key: str, + obs_date: date, + requested_asof: date, + has_pit: bool, + ) -> date: + if not has_pit: + return requested_asof + vintages = self._revision_map.get((series_key, obs_date)) + if vintages is None: + return requested_asof # no revision history → fall through + idx = min(self._n, len(vintages)) - 1 # 0-indexed + return vintages[idx] diff --git a/alphaforge/pit/views.py b/alphaforge/pit/views.py new file mode 100644 index 0000000..77a6e59 --- /dev/null +++ b/alphaforge/pit/views.py @@ -0,0 +1,51 @@ +"""Vintage view declarations for point-in-time data projection. + +A :class:`VintageView` is a pure value object that declares *what* vintage +resolution strategy should be applied — it carries no behavior. Pass it to a +:class:`~alphaforge.pit.resolvers.VintageResolver` to get resolution logic. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Literal + + +@dataclass(frozen=True) +class VintageView: + """Declares how the vintage dimension should be projected. + + This is a pure declaration with no behavior — it describes *what* view + is requested, not *how* to resolve it. Pass it to a VintageResolver to + get resolution behavior. + + Attributes: + mode: The vintage projection mode. + n_releases: Only used when ``mode="frozen"``. Which release to + freeze to (1-indexed). Defaults to 3, matching the common + BEA advance → second → third GDP release sequence. + """ + + mode: Literal["realtime", "latest", "frozen"] + n_releases: int = 3 # only used when mode="frozen" + + # Convenience constructors ------------------------------------------- + + @classmethod + def realtime(cls) -> VintageView: + """Identity view — return values as-of the requested date.""" + return cls(mode="realtime") + + @classmethod + def latest(cls) -> VintageView: + """Collapse all vintages to the most recent available.""" + return cls(mode="latest") + + @classmethod + def frozen(cls, n: int = 3) -> VintageView: + """Freeze each observation to its *n*-th release vintage. + + Args: + n: Release ordinal (1-indexed). Defaults to 3. + """ + return cls(mode="frozen", n_releases=n) diff --git a/alphaforge/pit/vintage.py b/alphaforge/pit/vintage.py new file mode 100644 index 0000000..873029a --- /dev/null +++ b/alphaforge/pit/vintage.py @@ -0,0 +1,52 @@ +"""Vintage selection and lookahead validation utilities. + +Moved from nowcast-data to alphaforge as generic PIT infrastructure. +""" +from __future__ import annotations + +from datetime import date, datetime + +import pandas as pd + + +def select_vintage_for_asof( + vintages: list[date | datetime | pd.Timestamp], + asof_date: date | datetime | pd.Timestamp, +) -> date | datetime | pd.Timestamp | None: + """Select the latest vintage not after asof_date. + + Returns None if no vintage is available before asof_date. + """ + if not vintages: + return None + + norm_vintages = [_normalize_date(v) for v in vintages] + norm_asof = _normalize_date(asof_date) + + sorted_vintages = sorted(zip(norm_vintages, vintages)) + + selected_original = None + for norm_v, orig_v in sorted_vintages: + if norm_v <= norm_asof: + selected_original = orig_v + else: + break + + return selected_original + + +def validate_no_lookahead( + vintage_date: date | datetime | pd.Timestamp, + asof_date: date | datetime | pd.Timestamp, +) -> bool: + """Return True if vintage_date <= asof_date (no lookahead bias).""" + return _normalize_date(vintage_date) <= _normalize_date(asof_date) + + +def _normalize_date(d: date | datetime | pd.Timestamp) -> pd.Timestamp: + """Normalize date to pd.Timestamp for comparison.""" + if isinstance(d, pd.Timestamp): + if d.tzinfo is not None: + return d.tz_localize(None) + return d + return pd.Timestamp(d) diff --git a/alphaforge/registry.py b/alphaforge/registry.py new file mode 100644 index 0000000..4c5584f --- /dev/null +++ b/alphaforge/registry.py @@ -0,0 +1,48 @@ +"""Generic entity registry mapping entity names to metadata.""" +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any + + +@dataclass(frozen=True) +class EntityEntry: + """Generic entity metadata entry.""" + + entity_id: str + asset_class: str + source_keys: dict[str, str] = field(default_factory=dict) # source → key pattern + metadata: dict[str, Any] = field(default_factory=dict) + + +class EntityRegistry: + """Generic registry mapping entity names to metadata.""" + + def __init__(self) -> None: + self._entries: dict[str, EntityEntry] = {} + + def register(self, name: str, entry: EntityEntry) -> None: + self._entries[name] = entry + + def get(self, name: str) -> EntityEntry: + if name not in self._entries: + raise KeyError(f"Unknown entity: {name!r}") + return self._entries[name] + + def entities(self, asset_class: str | None = None) -> list[str]: + if asset_class is None: + return sorted(self._entries) + return sorted(k for k, v in self._entries.items() if v.asset_class == asset_class) + + def source_key(self, entity: str, source: str) -> str: + entry = self.get(entity) + if source not in entry.source_keys: + raise ValueError(f"Entity {entity!r} has no key for source {source!r}") + return entry.source_keys[source] + + def all_source_keys(self, source: str) -> dict[str, str]: + result: dict[str, str] = {} + for name, entry in self._entries.items(): + if source in entry.source_keys: + result[name] = entry.source_keys[source] + return result diff --git a/docs/api/evaluation-metrics.md b/docs/api/evaluation-metrics.md new file mode 100644 index 0000000..d03a306 --- /dev/null +++ b/docs/api/evaluation-metrics.md @@ -0,0 +1,44 @@ +# Evaluation Metrics + +Pluggable forecast accuracy metrics for composing evaluation pipelines. + +The [`MetricFn`][alphaforge.evaluation.metrics.MetricFn] protocol defines +the interface that all metrics must satisfy. Five built-in implementations +are provided, and two convenience suites (`DEFAULT_METRICS` and +`BENCHMARK_METRICS`) bundle the most commonly used combinations. + +## Quick start + +```python +from alphaforge.evaluation.metrics import RMSE, MAE, BENCHMARK_METRICS + +# Single metric +rmse = RMSE() +score = rmse(y_pred, y_true) + +# Full benchmark suite +for metric in BENCHMARK_METRICS: + print(f"{metric.name}: {metric(y_pred, y_true):.4f}") +``` + +## Custom metrics + +Any class with a `name` attribute and `__call__(y_pred, y_true) -> float` +satisfies the protocol: + +```python +import numpy as np +from alphaforge.evaluation.metrics import MetricFn + +class MedianAbsoluteError: + name = "median_ae" + + def __call__(self, y_pred, y_true): + return float(np.median(np.abs(y_pred - y_true))) + +assert isinstance(MedianAbsoluteError(), MetricFn) +``` + +## API Reference + +::: alphaforge.evaluation.metrics diff --git a/docs/api/pit-missingness.md b/docs/api/pit-missingness.md new file mode 100644 index 0000000..ea3438c --- /dev/null +++ b/docs/api/pit-missingness.md @@ -0,0 +1,3 @@ +# PIT Missingness + +::: alphaforge.pit.missingness diff --git a/docs/api/pit-panel.md b/docs/api/pit-panel.md new file mode 100644 index 0000000..d8a8c35 --- /dev/null +++ b/docs/api/pit-panel.md @@ -0,0 +1,3 @@ +# PIT Panel Builder + +::: alphaforge.pit.panel diff --git a/docs/api/pit-release-rules.md b/docs/api/pit-release-rules.md new file mode 100644 index 0000000..df1b9a3 --- /dev/null +++ b/docs/api/pit-release-rules.md @@ -0,0 +1,3 @@ +# PIT Release Rules + +::: alphaforge.pit.release_rules diff --git a/docs/api/pit-resolvers.md b/docs/api/pit-resolvers.md new file mode 100644 index 0000000..66a1e9b --- /dev/null +++ b/docs/api/pit-resolvers.md @@ -0,0 +1,3 @@ +# PIT Vintage Resolvers + +::: alphaforge.pit.resolvers diff --git a/docs/api/pit-views.md b/docs/api/pit-views.md new file mode 100644 index 0000000..4e56467 --- /dev/null +++ b/docs/api/pit-views.md @@ -0,0 +1,3 @@ +# PIT Vintage Views + +::: alphaforge.pit.views diff --git a/docs/api/pit-vintage.md b/docs/api/pit-vintage.md new file mode 100644 index 0000000..efaa6e3 --- /dev/null +++ b/docs/api/pit-vintage.md @@ -0,0 +1,3 @@ +# PIT Vintage Utilities + +::: alphaforge.pit.vintage diff --git a/examples/short_rate_kim_dataset.py b/examples/short_rate_kim_dataset.py index 887f092..a719db3 100644 --- a/examples/short_rate_kim_dataset.py +++ b/examples/short_rate_kim_dataset.py @@ -10,8 +10,8 @@ from alphaforge import ( DataContext, DuckDBParquetStore, - FREDDataSource, FRBTermStructureBenchmarkSource, + FREDDataSource, PhiladelphiaSPFMeanLevelSource, TradingCalendar, build_kim_orphanides_dataset, diff --git a/mkdocs.yml b/mkdocs.yml index ff5068a..1d79985 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -47,10 +47,17 @@ nav: - Development: guides/development.md - API: - Package: api/package.md + - Evaluation Metrics: api/evaluation-metrics.md - DataContext: api/data-context.md - Dataset Builder: api/dataset-builder.md - Dataset Spec: api/dataset-spec.md - PIT Accessor: api/pit-accessor.md + - PIT Release Rules: api/pit-release-rules.md + - PIT Vintage Views: api/pit-views.md + - PIT Vintage Resolvers: api/pit-resolvers.md + - PIT Panel Builder: api/pit-panel.md + - PIT Missingness: api/pit-missingness.md + - PIT Vintage Utilities: api/pit-vintage.md - PIT Transforms: api/pit-transforms.md - PIT Pipelines: api/pit-pipelines.md - PIT Models: api/pit-models.md diff --git a/pyproject.toml b/pyproject.toml index 5e4a811..b3aa75d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -19,6 +19,7 @@ dependencies = [ "pyarrow>=14.0.1", "fredapi>=0.5.2", "PyYAML>=6.0.1", + "structlog>=24.0", ] [project.optional-dependencies] diff --git a/tests/fixtures/public_web/cftc_cot/sample.csv b/tests/fixtures/public_web/cftc_cot/sample.csv index e6d812e..66479ba 100644 --- a/tests/fixtures/public_web/cftc_cot/sample.csv +++ b/tests/fixtures/public_web/cftc_cot/sample.csv @@ -3,3 +3,7 @@ VIX FUTURES - CBOE FUTURES EXCHANGE,2026-01-06,1170E1,350000,45000,80000,12000,5 VIX FUTURES - CBOE FUTURES EXCHANGE,2026-01-13,1170E1,355000,48000,78000,13000,3000,-2000,92000,52000,21000,2000,2000,62000,68000,16000,2000,-2000,31000,26000,8500,1000,1000 VIX FUTURES - CBOE FUTURES EXCHANGE,2026-01-20,1170E1,360000,50000,75000,14000,2000,-3000,95000,55000,22000,3000,3000,65000,66000,17000,3000,-2000,32000,27000,9000,1000,1000 S&P 500 STOCK INDEX - CHICAGO MERCANTILE EXCHANGE,2026-01-06,13874+,500000,120000,90000,30000,10000,-5000,200000,150000,50000,5000,3000,100000,180000,40000,-2000,8000,50000,40000,10000,2000,-1000 +EURO FX - CHICAGO MERCANTILE EXCHANGE,2026-01-06,099741,420000,85000,65000,18000,4000,-2000,110000,95000,25000,3000,2000,75000,90000,20000,-500,3000,28000,22000,7000,800,-400 +EURO FX - CHICAGO MERCANTILE EXCHANGE,2026-01-13,099741,425000,88000,63000,19000,3000,-2000,112000,97000,26000,2000,2000,77000,88000,21000,2000,-2000,29000,23000,7500,1000,1000 +BRITISH POUND - CHICAGO MERCANTILE EXCHANGE,2026-01-06,096742,180000,35000,50000,8000,2000,-1500,55000,40000,12000,1500,1000,30000,45000,10000,-800,2000,15000,12000,4000,400,-300 +JAPANESE YEN - CHICAGO MERCANTILE EXCHANGE,2026-01-06,097741,250000,60000,40000,15000,3000,-1000,80000,70000,18000,2500,1800,50000,65000,14000,-300,2500,20000,18000,6000,600,-500 diff --git a/tests/fixtures/public_web/dtcc_ppd/sample_events.csv b/tests/fixtures/public_web/dtcc_ppd/sample_events.csv index bdd575d..1942823 100644 --- a/tests/fixtures/public_web/dtcc_ppd/sample_events.csv +++ b/tests/fixtures/public_web/dtcc_ppd/sample_events.csv @@ -1,4 +1,27 @@ execution_timestamp,reported_at_utc,asset_class,product,currency,tenor,price,notional,action,trade_id,venue,effective_date,maturity_date,cleared -2026-01-02T14:30:00Z,2026-01-02T14:35:00Z,Rates,IRS,USD,5Y,0.032,5000000,NEW,t1,SEF-A,2026-01-05,2031-01-05,TRUE -2026-01-02T15:10:00Z,2026-01-02T15:12:00Z,Rates,IRS,USD,5Y,0.034,2000000,NEW,t2,SEF-A,2026-01-05,2031-01-05,TRUE -2026-01-03T09:00:00Z,2026-01-03T09:02:00Z,Rates,OIS,EUR,2Y,0.018,800000,NEW,t3,SEF-B,2026-01-06,2028-01-06,FALSE +2026-01-02T14:30:00Z,2026-01-02T14:35:00Z,Rates,Interest Rate Swap,USD,5Y,0.032,5000000,NEW,t1,SEF-A,2026-01-05,2031-01-05,TRUE +2026-01-02T15:10:00Z,2026-01-02T15:12:00Z,Rates,Interest Rate Swap,USD,5Y,0.034,2000000,NEW,t2,SEF-A,2026-01-05,2031-01-05,TRUE +2026-01-02T15:45:00Z,2026-01-02T15:47:00Z,Rates,Interest Rate Swap,USD,10Y,0.038,8000000,NEW,t3,SEF-A,2026-01-05,2036-01-05,TRUE +2026-01-02T16:00:00Z,2026-01-02T16:02:00Z,Rates,Interest Rate Swap,USD,10Y,0.039,12000000,NEW,t4,SEF-B,2026-01-05,2036-01-05,TRUE +2026-01-02T16:30:00Z,2026-01-02T16:32:00Z,Rates,Interest Rate Swap,EUR,5Y,0.028,4000000,NEW,t5,SEF-A,2026-01-05,2031-01-05,TRUE +2026-01-02T17:00:00Z,2026-01-02T17:02:00Z,Rates,Interest Rate Swap,EUR,10Y,0.031,6000000,NEW,t6,SEF-B,2026-01-05,2036-01-05,TRUE +2026-01-02T17:15:00Z,2026-01-02T17:17:00Z,Rates,Interest Rate Swap,GBP,5Y,0.042,3000000,NEW,t7,SEF-A,2026-01-05,2031-01-05,TRUE +2026-01-02T17:30:00Z,2026-01-02T17:32:00Z,Rates,Interest Rate Swap,JPY,2Y,0.005,900000000,NEW,t8,SEF-B,2026-01-05,2028-01-05,TRUE +2026-01-03T09:00:00Z,2026-01-03T09:02:00Z,Rates,OIS,EUR,2Y,0.018,800000,NEW,t9,SEF-B,2026-01-06,2028-01-06,FALSE +2026-01-02T18:00:00Z,2026-01-02T18:02:00Z,Rates,Cross Currency Swap,USD,5Y,0.015,7000000,NEW,t10,SEF-A,2026-01-05,2031-01-05,TRUE +2026-01-02T18:15:00Z,2026-01-02T18:17:00Z,Rates,Cross Currency Swap,EUR,10Y,0.012,5000000,NEW,t11,SEF-B,2026-01-05,2036-01-05,TRUE +2026-01-02T18:30:00Z,2026-01-02T18:32:00Z,Rates,Cross Currency Swap,GBP,5Y,0.018,4000000,NEW,t12,SEF-A,2026-01-05,2031-01-05,TRUE +2026-01-02T18:45:00Z,2026-01-02T18:47:00Z,Rates,Cross Currency Swap,CHF,10Y,0.008,3000000,NEW,t13,SEF-B,2026-01-05,2036-01-05,TRUE +2026-01-02T19:00:00Z,2026-01-02T19:02:00Z,Rates,Cross Currency Swap,JPY,2Y,0.003,500000000,NEW,t14,SEF-A,2026-01-05,2028-01-05,TRUE +2026-01-02T10:00:00Z,2026-01-02T10:02:00Z,FX,FX Forward,USD,1M,1.0850,10000000,NEW,t15,,2026-02-05,,FALSE +2026-01-02T10:15:00Z,2026-01-02T10:17:00Z,FX,FX Forward,EUR,3M,1.0920,8000000,NEW,t16,,2026-04-05,,FALSE +2026-01-02T10:30:00Z,2026-01-02T10:32:00Z,FX,FX Forward,GBP,1M,1.2650,5000000,NEW,t17,,2026-02-05,,FALSE +2026-01-02T10:45:00Z,2026-01-02T10:47:00Z,FX,FX Forward,JPY,3M,148.50,15000000,NEW,t18,,2026-04-05,,FALSE +2026-01-02T11:00:00Z,2026-01-02T11:02:00Z,FX,FX Forward,CHF,6M,0.8900,6000000,NEW,t19,,2026-07-05,,FALSE +2026-01-02T11:30:00Z,2026-01-02T11:32:00Z,FX,FX Swap,USD,1M,0.0012,20000000,NEW,t20,,2026-02-05,,TRUE +2026-01-02T11:45:00Z,2026-01-02T11:47:00Z,FX,FX Swap,EUR,3M,0.0025,15000000,NEW,t21,,2026-04-05,,TRUE +2026-01-02T12:00:00Z,2026-01-02T12:02:00Z,FX,FX Swap,GBP,6M,0.0040,12000000,NEW,t22,,2026-07-05,,TRUE +2026-01-02T12:15:00Z,2026-01-02T12:17:00Z,FX,FX Swap,JPY,1M,0.0005,25000000,NEW,t23,,2026-02-05,,TRUE +2026-01-02T12:30:00Z,2026-01-02T12:32:00Z,FX,FX Swap,CHF,3M,0.0018,9000000,NEW,t24,,2026-04-05,,FALSE +2026-01-02T13:00:00Z,2026-01-02T13:02:00Z,Rates,Interest Rate Swap,USD,2Y,0.041,3000000,NEW,t25,SEF-A,2026-01-05,2028-01-05,TRUE +2026-01-02T13:15:00Z,2026-01-02T13:17:00Z,Rates,Interest Rate Swap,CHF,10Y,0.015,2000000,NEW,t26,SEF-B,2026-01-05,2036-01-05,TRUE diff --git a/tests/public_web/test_dtcc_ppd.py b/tests/public_web/test_dtcc_ppd.py index 38e1edf..9d9a4f9 100644 --- a/tests/public_web/test_dtcc_ppd.py +++ b/tests/public_web/test_dtcc_ppd.py @@ -1,10 +1,12 @@ from __future__ import annotations import io +import re import zipfile from pathlib import Path import pandas as pd +import pytest from alphaforge.data.public_web.dtcc_ppd import DTCCPPDSource from alphaforge.data.query import Query @@ -23,36 +25,77 @@ def _zip_fixture_bytes() -> bytes: return buf.getvalue() -def test_dtcc_events_and_daily_aggregation() -> None: - payload = _zip_fixture_bytes() +_FILE_ENTRY = { + "fileName": "CFTC_SLICE_RATES_2026_01_02_1.zip", + "dissemDTM": "2026-01-02T15:37:19Z", +} +_FX_FILE_ENTRY = { + "fileName": "CFTC_SLICE_FX_2026_01_02_1.zip", + "dissemDTM": "2026-01-02T10:37:19Z", +} - def list_provider(report_type: str, asset_code: str) -> list[dict]: - if report_type != "slice" or asset_code != "IR": - return [] - return [ - { - "fileName": "CFTC_SLICE_RATES_2026_01_02_1.zip", - "dissemDTM": "2026-01-02T15:37:19Z", - } - ] - def provider(file_name: str) -> bytes: - return payload +def _list_provider(report_type: str, asset_code: str) -> list[dict]: + if report_type == "slice" and asset_code == "IR": + return [_FILE_ENTRY] + if report_type == "slice" and asset_code == "FX": + return [_FX_FILE_ENTRY] + if report_type == "cumulative" and asset_code == "IR": + return [_FILE_ENTRY] + if report_type == "cumulative" and asset_code == "FX": + return [_FX_FILE_ENTRY] + return [] - source = DTCCPPDSource( - list_provider=list_provider, - artifact_provider=provider, - source_mode="slice", + +def _make_source(**kwargs) -> DTCCPPDSource: + payload = _zip_fixture_bytes() + defaults = dict( + list_provider=_list_provider, + artifact_provider=lambda fn: payload, + asset_codes=("IR", "FX"), ) + defaults.update(kwargs) + return DTCCPPDSource(**defaults) - events = source.fetch( - Query( - table="dtcc.ppd.events", - columns=["price", "notional", "product", "currency", "tenor"], - start=pd.Timestamp("2026-01-02", tz="UTC"), - end=pd.Timestamp("2026-01-02 23:59:59", tz="UTC"), - ) + +def _fetch_events(source: DTCCPPDSource | None = None, **query_kw) -> pd.DataFrame: + source = source or _make_source(source_mode="slice") + defaults = dict( + table="dtcc.ppd.events", + columns=["price", "notional", "product", "currency", "tenor", "asset_class"], + start=pd.Timestamp("2026-01-02", tz="UTC"), + end=pd.Timestamp("2026-01-03 23:59:59", tz="UTC"), + ) + defaults.update(query_kw) + return source.fetch(Query(**defaults)) + + +def _fetch_daily(source: DTCCPPDSource | None = None, **query_kw) -> pd.DataFrame: + source = source or _make_source(source_mode="cumulative") + defaults = dict( + table="dtcc.ppd.daily", + columns=[ + "trade_count", + "notional_sum", + "price_mean", + "price_std", + "notional_median", + "trade_count_large", + "dv01_proxy_sum", + ], + start=pd.Timestamp("2026-01-02", tz="UTC"), + end=pd.Timestamp("2026-01-03", tz="UTC"), ) + defaults.update(query_kw) + return source.fetch(Query(**defaults)) + + +# ── Original test (updated for expanded fixture) ────────────────────── + + +def test_dtcc_events_and_daily_aggregation() -> None: + source = _make_source(source_mode="slice") + events = _fetch_events(source) assert not events.empty assert str(events["ts_utc"].dtype).startswith(("datetime64[ns,", "datetime64[us,")) @@ -75,3 +118,154 @@ def provider(file_name: str) -> bytes: assert int(daily["trade_count"].iloc[0]) >= 1 assert "notional_sum" in daily.columns assert str(daily["asof_utc"].dtype).startswith(("datetime64[ns,", "datetime64[us,")) + + +# ── New tests ───────────────────────────────────────────────────────── + + +class TestSchemas: + def test_schemas(self) -> None: + source = _make_source() + schemas = source.schemas() + assert "dtcc.ppd.events" in schemas + assert "dtcc.ppd.daily" in schemas + + ev = schemas["dtcc.ppd.events"] + assert ev.time_column == "ts_utc" + assert ev.entity_column == "entity_id" + assert ev.native_freq == "D" + assert ev.time_semantics == "point" + + daily = schemas["dtcc.ppd.daily"] + assert daily.time_column == "date" + assert daily.entity_column == "entity_id" + assert daily.native_freq == "D" + assert daily.time_semantics == "interval_end" + + +class TestEntityIdFormat: + def test_entity_id_format(self) -> None: + events = _fetch_events() + pattern = re.compile(r"^dtccppd\.[a-z_]+\.[a-z_]+\.[a-z]+\.\w+$") + for eid in events["entity_id"].unique(): + assert pattern.match(eid), f"Entity ID does not match pattern: {eid}" + assert eid == eid.lower(), f"Entity ID has uppercase: {eid}" + assert " " not in eid, f"Entity ID has spaces: {eid}" + + +class TestFXEvents: + def test_fx_events(self) -> None: + events = _fetch_events() + fx = events[events["entity_id"].str.startswith("dtccppd.fx.")] + assert not fx.empty, "No FX events found" + products = {p.lower().replace(" ", "_") for p in fx["product"].unique()} + assert "fx_forward" in products or "fx_swap" in products, ( + f"Expected fx_forward or fx_swap in FX products, got {products}" + ) + + +class TestIRSEvents: + def test_irs_events(self) -> None: + events = _fetch_events() + irs = events[events["entity_id"].str.contains("interest_rate_swap")] + assert not irs.empty, "No IRS events found" + + +class TestCCSEvents: + def test_ccs_events(self) -> None: + events = _fetch_events() + ccs = events[events["entity_id"].str.contains("cross_currency_swap")] + assert not ccs.empty, "No CCS events found" + + +class TestDailyAggregation: + def test_daily_aggregation_columns(self) -> None: + daily = _fetch_daily() + assert not daily.empty + required = [ + "trade_count", + "notional_sum", + "price_mean", + "price_std", + "notional_median", + "trade_count_large", + "dv01_proxy_sum", + ] + for col in required: + assert col in daily.columns, f"Missing column: {col}" + + +class TestDailyDV01Proxy: + def test_daily_dv01_proxy(self) -> None: + # Use IR-only source to avoid double-counting from FX payload + source = _make_source(source_mode="cumulative", asset_codes=("IR",)) + daily = _fetch_daily(source) + # Find a 10y IRS entity + irs_10y = daily[ + daily["entity_id"].str.contains("interest_rate_swap") + & daily["entity_id"].str.contains("10y") + ] + assert not irs_10y.empty, "No 10y IRS entity in daily" + + # Compute expected: from fixture, 10y IRS USD rows have + # notional 8_000_000 + 12_000_000 = 20_000_000 + # duration_scalar for 10y = 8.5 + # dv01_proxy per row = notional * 8.5 * 1e-4 + # total = (8_000_000 + 12_000_000) * 8.5 * 1e-4 = 17_000 + row = irs_10y[irs_10y["entity_id"].str.contains("usd")] + if not row.empty: + expected = (8_000_000 + 12_000_000) * 8.5 * 1e-4 + assert abs(row["dv01_proxy_sum"].iloc[0] - expected) < 1.0, ( + f"Expected dv01_proxy_sum ≈ {expected}, got {row['dv01_proxy_sum'].iloc[0]}" + ) + + +class TestEntityFilter: + def test_entity_filter(self) -> None: + events_all = _fetch_events() + target = events_all["entity_id"].iloc[0] + + source = _make_source(source_mode="slice") + filtered = source.fetch( + Query( + table="dtcc.ppd.events", + columns=["price", "notional"], + start=pd.Timestamp("2026-01-02", tz="UTC"), + end=pd.Timestamp("2026-01-03 23:59:59", tz="UTC"), + entities=[target], + ) + ) + assert not filtered.empty + assert (filtered["entity_id"] == target).all() + + +class TestTimeFilter: + def test_time_filter(self) -> None: + source = _make_source(source_mode="slice") + # Tight window: only 14:00-15:00 on 2026-01-02 + filtered = source.fetch( + Query( + table="dtcc.ppd.events", + columns=["price", "notional"], + start=pd.Timestamp("2026-01-02T14:00:00Z"), + end=pd.Timestamp("2026-01-02T15:00:00Z"), + ) + ) + if not filtered.empty: + ts = pd.DatetimeIndex(filtered["ts_utc"]) + assert (ts >= pd.Timestamp("2026-01-02T14:00:00Z")).all() + assert (ts <= pd.Timestamp("2026-01-02T15:00:00Z")).all() + + +class TestUnknownTable: + def test_unknown_table_raises(self) -> None: + source = _make_source() + with pytest.raises(ValueError, match="Unknown table"): + source.fetch( + Query( + table="dtcc.ppd.bad_table", + columns=["price"], + start=pd.Timestamp("2026-01-02", tz="UTC"), + end=pd.Timestamp("2026-01-03", tz="UTC"), + ) + ) diff --git a/tests/test_config.py b/tests/test_config.py new file mode 100644 index 0000000..8423f3a --- /dev/null +++ b/tests/test_config.py @@ -0,0 +1,26 @@ +"""Tests for alphaforge config resolution.""" +from __future__ import annotations + +from alphaforge.config import ConfigEntry, resolve_config_value + + +class TestResolveConfigValue: + def test_explicit_wins(self) -> None: + assert resolve_config_value("/explicit", "TEST_ENV", "/default") == "/explicit" + + def test_env_var_used(self, monkeypatch) -> None: + monkeypatch.setenv("TEST_CFG_VAR", "/from_env") + assert resolve_config_value(None, "TEST_CFG_VAR", "/default") == "/from_env" + + def test_default_fallback(self) -> None: + assert resolve_config_value(None, "UNLIKELY_ENV_VAR_XYZ", "/default") == "/default" + + +class TestConfigEntry: + def test_resolve_explicit(self) -> None: + entry = ConfigEntry(name="store", env_var="TEST_STORE", default="/def") + assert entry.resolve("/override") == "/override" + + def test_resolve_default(self) -> None: + entry = ConfigEntry(name="store", env_var="UNLIKELY_ENV_XYZ", default="/def") + assert entry.resolve() == "/def" diff --git a/tests/test_logging.py b/tests/test_logging.py new file mode 100644 index 0000000..e952748 --- /dev/null +++ b/tests/test_logging.py @@ -0,0 +1,73 @@ +"""Tests for alphaforge.logging module.""" +from __future__ import annotations + +import logging +from pathlib import Path + +from structlog.testing import capture_logs + +from alphaforge.logging import configure_logging, get_logger + + +class TestConfigureConsole: + def test_configure_console(self) -> None: + configure_logging(level="DEBUG", format="console") + log = get_logger("test_console") + log.info("console test") # should not crash + + +class TestConfigureJson: + def test_configure_json(self, capsys) -> None: + configure_logging(level="DEBUG", format="json") + log = get_logger("test_json") + log.info("json test", key="value") + + +class TestGetLogger: + def test_get_logger(self) -> None: + log = get_logger("my_module") + assert log is not None + + +class TestBoundContext: + def test_bound_context(self) -> None: + configure_logging(level="DEBUG", format="console") + with capture_logs() as cap: + log = get_logger("test_bound") + log.info("hello", source="test") + events = [e for e in cap if e.get("event") == "hello"] + assert len(events) == 1 + assert events[0]["source"] == "test" + + +class TestFileLogging: + def test_file_logging(self, tmp_path: Path) -> None: + log_file = tmp_path / "test.log" + configure_logging(level="DEBUG", format="console", log_file=str(log_file)) + log = get_logger("test_file") + log.info("file test", data=42) + # File should have been written + assert log_file.exists() + content = log_file.read_text() + assert "file test" in content + + +class TestCaptureLogs: + def test_capture_logs(self) -> None: + configure_logging(level="DEBUG", format="console") + with capture_logs() as cap: + log = get_logger("test_capture") + log.info("captured", rows=10) + events = [e for e in cap if e.get("event") == "captured"] + assert len(events) == 1 + assert events[0]["rows"] == 10 + + +class TestStdlibIntegration: + def test_stdlib_integration(self, capsys) -> None: + configure_logging(level="DEBUG", format="console") + stdlib_logger = logging.getLogger("stdlib_test") + stdlib_logger.info("stdlib message") + # stdlib messages go through structlog's formatter to stderr + captured = capsys.readouterr() + assert "stdlib message" in captured.err diff --git a/tests/test_missingness.py b/tests/test_missingness.py new file mode 100644 index 0000000..d6a8667 --- /dev/null +++ b/tests/test_missingness.py @@ -0,0 +1,37 @@ +"""Tests for missingness classification.""" +from __future__ import annotations + +from datetime import date + +from alphaforge.pit.missingness import MissingnessReason, classify_missingness +from alphaforge.pit.release_rules import FixedLagMonths + + +class TestMissingnessReason: + def test_structural_quarterly_in_monthly(self) -> None: + result = classify_missingness( + obs_date=date(2025, 2, 28), + asof_date=date(2025, 4, 1), + series_frequency="Q", + panel_frequency="M", + ) + assert result == MissingnessReason.STRUCTURAL + + def test_future_obs(self) -> None: + result = classify_missingness( + obs_date=date(2025, 6, 30), + asof_date=date(2025, 5, 15), + series_frequency="M", + ) + assert result == MissingnessReason.FUTURE + + def test_ragged_edge_with_release_rule(self) -> None: + # obs_date=Mar 31, release expected May 1 (2 month lag), asof=Apr 15 (before release) + rule = FixedLagMonths(months=2) + result = classify_missingness( + obs_date=date(2025, 3, 31), + asof_date=date(2025, 4, 15), + series_frequency="M", + release_rule=rule, + ) + assert result == MissingnessReason.RAGGED_EDGE diff --git a/tests/test_pipeline_health.py b/tests/test_pipeline_health.py new file mode 100644 index 0000000..de22af4 --- /dev/null +++ b/tests/test_pipeline_health.py @@ -0,0 +1,112 @@ +"""Tests for alphaforge.pipeline.health module.""" +from __future__ import annotations + +import pandas as pd + +from alphaforge.pipeline.health import ( + CFTC_COT_HEALTH_POLICY, + DTCC_PPD_HEALTH_POLICY, + SourceHealthPolicy, + assess_source_health, +) + + +def _ts(days_ago: float) -> pd.Timestamp: + return pd.Timestamp.now(tz="UTC") - pd.Timedelta(days=days_ago) + + +_POLICY = SourceHealthPolicy( + expected_cadence=pd.Timedelta(days=7), + grace_period=pd.Timedelta(days=2), + stale_threshold=pd.Timedelta(days=14), + dead_threshold=pd.Timedelta(days=42), + weight_decay_half_life=pd.Timedelta(days=7), +) +_NOW = pd.Timestamp.now(tz="UTC") + + +class TestOkStatus: + def test_ok_status(self) -> None: + h = assess_source_health("test", _ts(5), _NOW, _POLICY) + assert h.status == "ok" + assert h.weight_factor == 1.0 + + +class TestLateStatus: + def test_late_status(self) -> None: + # Past cadence+grace (9d) but within stale (14d) + h = assess_source_health("test", _ts(12), _NOW, _POLICY) + assert h.status == "late" + assert h.weight_factor == 1.0 + + +class TestStaleStatus: + def test_stale_status(self) -> None: + h = assess_source_health("test", _ts(21), _NOW, _POLICY) + assert h.status == "stale" + assert 0 < h.weight_factor < 1.0 + + +class TestStaleWeightDecay: + def test_stale_weight_decay(self) -> None: + # At stale_threshold (14d) + 1 half-life (7d) = 21d → weight ≈ 0.5 + h = assess_source_health("test", _ts(21), _NOW, _POLICY) + assert abs(h.weight_factor - 0.5) < 0.05 + + # At stale_threshold + 2 half-lives = 28d → weight ≈ 0.25 + h2 = assess_source_health("test", _ts(28), _NOW, _POLICY) + assert abs(h2.weight_factor - 0.25) < 0.05 + + +class TestDeadStatus: + def test_dead_status(self) -> None: + h = assess_source_health("test", _ts(50), _NOW, _POLICY) + assert h.status == "dead" + assert h.weight_factor == 0.0 + + +class TestEmptyStatus: + def test_empty_status(self) -> None: + h = assess_source_health("test", None, _NOW, _POLICY) + assert h.status == "empty" + assert h.weight_factor == 0.0 + + +class TestCftcPolicyShutdown: + def test_cftc_policy_shutdown(self) -> None: + """Simulate a 35-day gov shutdown (2018 style).""" + policy = CFTC_COT_HEALTH_POLICY + # Day 5: ok + h5 = assess_source_health("cftc", _ts(5), _NOW, policy) + assert h5.status == "ok" + # Day 15: late + h15 = assess_source_health("cftc", _ts(15), _NOW, policy) + assert h15.status == "late" + # Day 25: stale + h25 = assess_source_health("cftc", _ts(25), _NOW, policy) + assert h25.status == "stale" + assert h25.weight_factor > 0 + # Day 50: still stale (dead_threshold is 56) + h50 = assess_source_health("cftc", _ts(50), _NOW, policy) + assert h50.status == "stale" + # Day 60: dead + h60 = assess_source_health("cftc", _ts(60), _NOW, policy) + assert h60.status == "dead" + + +class TestDtccPolicyWeekend: + def test_dtcc_policy_weekend(self) -> None: + """Saturday/Sunday should be ok due to grace period.""" + policy = DTCC_PPD_HEALTH_POLICY + # 2 days ago (e.g., Friday data on Sunday) + h = assess_source_health("dtcc", _ts(2), _NOW, policy) + assert h.status == "ok" + + +class TestExpectedNext: + def test_expected_next(self) -> None: + latest = _ts(3) + h = assess_source_health("test", latest, _NOW, _POLICY) + expected = latest + pd.Timedelta(days=7) + assert h.expected_next is not None + assert abs((h.expected_next - expected).total_seconds()) < 1 diff --git a/tests/test_pipeline_protocols.py b/tests/test_pipeline_protocols.py new file mode 100644 index 0000000..5cb8c17 --- /dev/null +++ b/tests/test_pipeline_protocols.py @@ -0,0 +1,158 @@ +"""Tests for alphaforge.pipeline protocols and SimplePipeline.""" +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +import pandas as pd +import pytest + +from alphaforge.pipeline.protocols import ( + Filter, + Parametric, + PipelineVariant, + Signal, + SimplePipeline, + Transformer, +) + +# --- Stub implementations --- + + +@dataclass +class StubFilter: + name: str = "stub_filter" + keep_above: float = 0.0 + + def apply(self, df: pd.DataFrame, **params: Any) -> pd.DataFrame: + return df[df["value"] > self.keep_above].copy() + + +@dataclass +class StubTransformer: + name: str = "stub_transformer" + multiplier: float = 2.0 + + def transform(self, df: pd.DataFrame, **params: Any) -> pd.DataFrame: + out = df.copy() + out["value"] = out["value"] * self.multiplier + return out + + def get_params(self) -> dict[str, float]: + return {"multiplier": self.multiplier} + + def set_params(self, params: dict[str, float]) -> None: + if "multiplier" in params: + self.multiplier = params["multiplier"] + + +@dataclass +class StubSignal: + name: str = "stub_signal" + + def score(self, df: pd.DataFrame, **params: Any) -> pd.Series: + return df.groupby("entity_id")["value"].last().rename(self.name) + + +@dataclass +class StubRemoveAllFilter: + name: str = "remove_all" + + def apply(self, df: pd.DataFrame, **params: Any) -> pd.DataFrame: + return df.iloc[0:0].copy() + + +def _make_df() -> pd.DataFrame: + return pd.DataFrame({ + "entity_id": ["a", "a", "b", "b", "c", "c"], + "date": pd.date_range("2026-01-01", periods=6, freq="D"), + "value": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0], + }) + + +# --- Tests --- + + +class TestFilterProtocol: + def test_filter_protocol(self) -> None: + f = StubFilter() + assert isinstance(f, Filter) + + +class TestTransformerProtocol: + def test_transformer_protocol(self) -> None: + t = StubTransformer() + assert isinstance(t, Transformer) + + +class TestSignalProtocol: + def test_signal_protocol(self) -> None: + s = StubSignal() + assert isinstance(s, Signal) + + +class TestParametricProtocol: + def test_parametric_protocol(self) -> None: + t = StubTransformer() + assert isinstance(t, Parametric) + + +class TestSimplePipelineChain: + def test_simple_pipeline_chain(self) -> None: + pipe = SimplePipeline( + name="test", + filters=[StubFilter(keep_above=2.0)], + transformers=[StubTransformer(multiplier=10.0)], + signal=StubSignal(), + ) + df = _make_df() + result = pipe.run(df) + assert not result.empty + # After filter: values > 2.0 → {3,4,5,6}, after mult×10 → {30,40,50,60} + # Last per entity: b→40, c→60 + assert result["b"] == 40.0 + assert result["c"] == 60.0 + + +class TestSimplePipelineNoSignalRaises: + def test_simple_pipeline_no_signal_raises(self) -> None: + pipe = SimplePipeline(name="no_signal") + with pytest.raises(ValueError, match="no terminal Signal"): + pipe.run(_make_df()) + + +class TestSimplePipelineGetSetParams: + def test_simple_pipeline_get_set_params(self) -> None: + t = StubTransformer(multiplier=5.0) + pipe = SimplePipeline(name="test", transformers=[t], signal=StubSignal()) + + params = pipe.get_all_params() + assert params == {"stub_transformer": {"multiplier": 5.0}} + + pipe.set_all_params({"stub_transformer": {"multiplier": 3.0}}) + assert t.multiplier == 3.0 + + +class TestSimplePipelineEmptyAfterFilter: + def test_simple_pipeline_empty_after_filter(self) -> None: + pipe = SimplePipeline( + name="test", + filters=[StubRemoveAllFilter()], + transformers=[StubTransformer()], + signal=StubSignal(), + ) + result = pipe.run(_make_df()) + assert result.empty + + +class TestPipelineVariant: + def test_pipeline_variant(self) -> None: + pipe = SimplePipeline( + name="inner", + transformers=[StubTransformer(multiplier=1.0)], + signal=StubSignal(), + ) + variant = PipelineVariant(name="v1", pipeline=pipe, description="test variant") + result = variant.pipeline.run(_make_df()) + assert not result.empty + assert variant.name == "v1" diff --git a/tests/test_pipeline_weights.py b/tests/test_pipeline_weights.py new file mode 100644 index 0000000..66480d5 --- /dev/null +++ b/tests/test_pipeline_weights.py @@ -0,0 +1,36 @@ +"""Tests for pipeline weight adjustment.""" +from __future__ import annotations + +from alphaforge.pipeline.weights import adjust_weights_for_health + + +class TestAdjustWeightsForHealth: + def test_no_degradation(self) -> None: + weights = {"a": 0.5, "b": 0.5} + result = adjust_weights_for_health( + weights=weights, + health_fn=lambda _: 1.0, + source_map={"a": "src_a", "b": "src_b"}, + ) + assert result == {"a": 0.5, "b": 0.5} + + def test_partial_degradation_renormalized(self) -> None: + weights = {"a": 0.5, "b": 0.5} + result = adjust_weights_for_health( + weights=weights, + health_fn=lambda s: 0.5 if s == "src_a" else 1.0, + source_map={"a": "src_a", "b": "src_b"}, + ) + # a: 0.5*0.5 = 0.25, b: 0.5*1.0 = 0.5, total = 0.75 + # renormalized: a = 0.25 * (1.0/0.75), b = 0.5 * (1.0/0.75) + assert abs(result["a"] + result["b"] - 1.0) < 1e-10 + + def test_dead_source_gets_zero(self) -> None: + weights = {"a": 0.5, "b": 0.5} + result = adjust_weights_for_health( + weights=weights, + health_fn=lambda s: 0.0 if s == "src_a" else 1.0, + source_map={"a": "src_a", "b": "src_b"}, + ) + assert result["a"] == 0.0 + assert abs(result["b"] - 1.0) < 1e-10 diff --git a/tests/test_pit_panel_builder.py b/tests/test_pit_panel_builder.py new file mode 100644 index 0000000..a1796f6 --- /dev/null +++ b/tests/test_pit_panel_builder.py @@ -0,0 +1,87 @@ +"""Tests for PIT panel builder and vintage utilities.""" +from __future__ import annotations + +from datetime import date + +import pandas as pd + +from alphaforge.pit.vintage import select_vintage_for_asof, validate_no_lookahead + + +class TestSelectVintageForAsof: + def test_exact_match(self) -> None: + vintages = [date(2025, 1, 1), date(2025, 2, 1), date(2025, 3, 1)] + result = select_vintage_for_asof(vintages, date(2025, 2, 1)) + assert result == date(2025, 2, 1) + + def test_between_vintages(self) -> None: + vintages = [date(2025, 1, 1), date(2025, 2, 1)] + result = select_vintage_for_asof(vintages, date(2025, 1, 15)) + assert result == date(2025, 1, 1) + + def test_before_all_vintages(self) -> None: + vintages = [date(2025, 2, 1)] + result = select_vintage_for_asof(vintages, date(2025, 1, 1)) + assert result is None + + def test_empty_vintages(self) -> None: + assert select_vintage_for_asof([], date(2025, 1, 1)) is None + + def test_after_all_vintages(self) -> None: + vintages = [date(2025, 1, 1), date(2025, 2, 1)] + result = select_vintage_for_asof(vintages, date(2025, 6, 1)) + assert result == date(2025, 2, 1) + + +class TestValidateNoLookahead: + def test_valid(self) -> None: + assert validate_no_lookahead(date(2025, 1, 1), date(2025, 2, 1)) is True + + def test_invalid(self) -> None: + assert validate_no_lookahead(date(2025, 3, 1), date(2025, 2, 1)) is False + + def test_equal(self) -> None: + assert validate_no_lookahead(date(2025, 1, 1), date(2025, 1, 1)) is True + + +class TestBuildPitPanel: + def test_build_pit_panel_basic(self, tmp_path) -> None: + from alphaforge.pit.accessor import PITAccessor, ensure_pit_table + from alphaforge.pit.panel import build_pit_panel + from alphaforge.store.duckdb_parquet import DuckDBParquetStore + + store = DuckDBParquetStore(root=str(tmp_path / "store")) + conn = store.conn() + ensure_pit_table(conn) + pit = PITAccessor(conn) + + # Insert some test data + obs = pd.DataFrame([ + {"series_key": "test.a", "obs_date": pd.Timestamp("2025-01-01"), "asof_utc": pd.Timestamp("2025-01-02", tz="UTC"), "value": 1.0, "source": "test"}, + {"series_key": "test.a", "obs_date": pd.Timestamp("2025-01-08"), "asof_utc": pd.Timestamp("2025-01-09", tz="UTC"), "value": 2.0, "source": "test"}, + {"series_key": "test.b", "obs_date": pd.Timestamp("2025-01-01"), "asof_utc": pd.Timestamp("2025-01-02", tz="UTC"), "value": 10.0, "source": "test"}, + ]) + pit.upsert_pit_observations(obs, strict="coerce") + + panel = build_pit_panel( + pit, + series_keys={"col_a": "test.a", "col_b": "test.b"}, + asof=pd.Timestamp("2025-02-01", tz="UTC"), + ) + assert "col_a" in panel.columns + assert "col_b" in panel.columns + assert len(panel) >= 1 + + +class TestLongToWide: + def test_basic_pivot(self) -> None: + from alphaforge.pit.panel import long_to_wide + + df = pd.DataFrame({ + "obs_date": [pd.Timestamp("2025-01-01")] * 2, + "series_key": ["a", "b"], + "value": [1.0, 2.0], + }) + wide = long_to_wide(df) + assert "a" in wide.columns + assert "b" in wide.columns diff --git a/tests/test_registry.py b/tests/test_registry.py new file mode 100644 index 0000000..ec38e47 --- /dev/null +++ b/tests/test_registry.py @@ -0,0 +1,35 @@ +"""Tests for the generic EntityRegistry.""" +from __future__ import annotations + +import pytest + +from alphaforge.registry import EntityEntry, EntityRegistry + + +class TestEntityRegistry: + def test_register_and_get(self) -> None: + reg = EntityRegistry() + entry = EntityEntry(entity_id="eur", asset_class="fx", source_keys={"cot": "k1"}) + reg.register("eur", entry) + assert reg.get("eur") is entry + + def test_get_unknown_raises(self) -> None: + reg = EntityRegistry() + with pytest.raises(KeyError, match="Unknown entity"): + reg.get("missing") + + def test_entities_filter_by_asset_class(self) -> None: + reg = EntityRegistry() + reg.register("eur", EntityEntry("eur", "fx")) + reg.register("usd_irs_5y", EntityEntry("usd_irs_5y", "rates")) + assert reg.entities("fx") == ["eur"] + assert reg.entities("rates") == ["usd_irs_5y"] + assert sorted(reg.entities()) == ["eur", "usd_irs_5y"] + + def test_source_key_and_all_source_keys(self) -> None: + reg = EntityRegistry() + reg.register("eur", EntityEntry("eur", "fx", source_keys={"cot": "cot_eur", "dtcc": "dtcc_eur"})) + reg.register("gbp", EntityEntry("gbp", "fx", source_keys={"cot": "cot_gbp"})) + assert reg.source_key("eur", "cot") == "cot_eur" + assert reg.all_source_keys("cot") == {"eur": "cot_eur", "gbp": "cot_gbp"} + assert reg.all_source_keys("dtcc") == {"eur": "dtcc_eur"} diff --git a/tests/test_release_rules.py b/tests/test_release_rules.py new file mode 100644 index 0000000..a7c555c --- /dev/null +++ b/tests/test_release_rules.py @@ -0,0 +1,79 @@ +"""Tests for release rule schedule computations.""" + +from datetime import date + +import pytest + +from alphaforge.pit.release_rules import ( + RULE_REGISTRY, + CustomRule, + NthBusinessDay, + QuarterlyRelease, + ReleaseRule, + WeeklyRelease, +) + + +class TestRuleRegistry: + def test_all_rules_registered(self): + expected = { + "nth_business_day", + "nth_weekday", + "calendar_day", + "fixed_lag_months", + "quarterly_release", + "weekly", + "custom", + } + assert set(RULE_REGISTRY.keys()) == expected + + def test_round_trip_from_dict(self): + rule = NthBusinessDay(n=3, anchor="following_month") + d = rule.to_dict() + assert d["type"] == "nth_business_day" + restored = ReleaseRule.from_dict(d) + assert isinstance(restored, NthBusinessDay) + assert restored.n == 3 + + def test_from_dict_unknown_type(self): + with pytest.raises(ValueError, match="Unknown release rule type"): + ReleaseRule.from_dict({"type": "nonexistent"}) + + +class TestNthBusinessDay: + def test_first_business_day_following_month(self): + rule = NthBusinessDay(n=1, anchor="following_month") + result = rule.expected_release_date(date(2025, 1, 31)) + assert result == date(2025, 2, 3) + + def test_third_business_day(self): + rule = NthBusinessDay(n=3, anchor="following_month") + result = rule.expected_release_date(date(2024, 12, 31)) + assert result == date(2025, 1, 6) + + +class TestQuarterlyRelease: + def test_advance_release(self): + rule = QuarterlyRelease(advance_lag_months=1, preliminary_lag_months=2, final_lag_months=3) + result = rule.expected_release_date(date(2024, 12, 31), release_number=1) + assert result == date(2025, 1, 1) + + def test_final_release(self): + rule = QuarterlyRelease(advance_lag_months=1, preliminary_lag_months=2, final_lag_months=3) + result = rule.expected_release_date(date(2024, 12, 31), release_number=3) + assert result == date(2025, 3, 1) + + +class TestWeeklyRelease: + def test_five_day_lag(self): + rule = WeeklyRelease(release_weekday="Thursday", lag_days=5) + result = rule.expected_release_date(date(2025, 1, 4)) + assert result == date(2025, 1, 9) + + +class TestCustomRule: + def test_serialization(self): + rule = CustomRule(description="test", approximate_lag_months=1) + d = rule.to_dict() + restored = ReleaseRule.from_dict(d) + assert isinstance(restored, CustomRule) diff --git a/tests/test_torch_protocols.py b/tests/test_torch_protocols.py new file mode 100644 index 0000000..d066280 --- /dev/null +++ b/tests/test_torch_protocols.py @@ -0,0 +1,110 @@ +"""Tests for alphaforge.pipeline.torch_protocols module.""" +from __future__ import annotations + +import numpy as np +import pytest + +torch = pytest.importorskip("torch") +nn = torch.nn + +from alphaforge.pipeline.torch_protocols import ( # noqa: E402 + DifferentiableSignalModule, + NumpyDifferentiableWrapper, + TorchParametric, +) + + +class _SimpleModule(DifferentiableSignalModule): + name = "simple" + + def __init__(self): + super().__init__() + self.scale = nn.Parameter(torch.tensor(2.0)) + + def forward(self, x): + return x * self.scale + + +class TestTorchParametricProtocol: + def test_torch_parametric_protocol(self) -> None: + m = _SimpleModule() + assert isinstance(m, TorchParametric) + + +class TestDifferentiableSignalModuleGetSet: + def test_get_set(self) -> None: + m = _SimpleModule() + params = m.get_params() + assert "scale" in params + assert abs(params["scale"] - 2.0) < 1e-6 + m.set_params({"scale": 5.0}) + assert abs(m.scale.item() - 5.0) < 1e-6 + + +class TestDifferentiableSignalModuleGradients: + def test_gradients(self) -> None: + m = _SimpleModule() + x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True) + y = m(x) + loss = y.sum() + loss.backward() + assert m.scale.grad is not None + assert x.grad is not None + + +class TestNumpyWrapperForward: + def test_numpy_wrapper_forward(self) -> None: + def fwd(x_np, params): + return x_np * params.get("scale", 1.0) + + wrapper = NumpyDifferentiableWrapper( + forward_fn=fwd, + param_names=["scale"], + init_params={"scale": 3.0}, + name="test_wrapper", + ) + x = torch.tensor([1.0, 2.0, 3.0]) + result = wrapper(x) + np.testing.assert_allclose(result.detach().numpy(), [3.0, 6.0, 9.0]) + + +class TestNumpyWrapperAnalyticalGrad: + def test_numpy_wrapper_analytical_grad(self) -> None: + def fwd(x_np, params): + return x_np * params["scale"] + + def bwd(x_np, grad_np, params): + return {"input": grad_np * params["scale"]} + + wrapper = NumpyDifferentiableWrapper( + forward_fn=fwd, + backward_fn=bwd, + param_names=["scale"], + init_params={"scale": 2.0}, + ) + x = torch.tensor([1.0, 2.0], requires_grad=True) + y = wrapper(x) + loss = y.sum() + loss.backward() + assert x.grad is not None + np.testing.assert_allclose(x.grad.numpy(), [2.0, 2.0]) + + +class TestNumpyWrapperFiniteDiff: + def test_numpy_wrapper_finite_diff(self) -> None: + def fwd(x_np, params): + return x_np ** 2 * params.get("scale", 1.0) + + wrapper = NumpyDifferentiableWrapper( + forward_fn=fwd, + param_names=["scale"], + init_params={"scale": 1.0}, + ) + x = torch.tensor([2.0, 3.0], requires_grad=True) + y = wrapper(x) + loss = y.sum() + loss.backward() + # Finite diff should give approximate gradients + assert x.grad is not None + # d/dx(x^2) = 2x at x=[2,3] → [4,6] + np.testing.assert_allclose(x.grad.numpy(), [4.0, 6.0], atol=0.1) diff --git a/tests/test_vintage_resolvers.py b/tests/test_vintage_resolvers.py new file mode 100644 index 0000000..1282eb2 --- /dev/null +++ b/tests/test_vintage_resolvers.py @@ -0,0 +1,253 @@ +"""Tests for vintage view resolvers.""" + +from __future__ import annotations + +from datetime import date + +import pytest + +from alphaforge.pit.resolvers import ( + FrozenResolver, + LatestResolver, + RealtimeResolver, + VintageResolver, +) +from alphaforge.pit.views import VintageView + +# --------------------------------------------------------------------------- +# VintageView value object +# --------------------------------------------------------------------------- + + +class TestVintageView: + def test_realtime_factory(self) -> None: + v = VintageView.realtime() + assert v.mode == "realtime" + assert v.n_releases == 3 # default, unused + + def test_latest_factory(self) -> None: + v = VintageView.latest() + assert v.mode == "latest" + + def test_frozen_factory_default(self) -> None: + v = VintageView.frozen() + assert v.mode == "frozen" + assert v.n_releases == 3 + + def test_frozen_factory_custom(self) -> None: + v = VintageView.frozen(n=5) + assert v.mode == "frozen" + assert v.n_releases == 5 + + def test_frozen_immutable(self) -> None: + v = VintageView.frozen() + with pytest.raises(AttributeError): + v.mode = "latest" # type: ignore[misc] + + def test_equality(self) -> None: + assert VintageView.realtime() == VintageView(mode="realtime") + assert VintageView.frozen(3) == VintageView.frozen(3) + assert VintageView.frozen(3) != VintageView.frozen(5) + + +# --------------------------------------------------------------------------- +# RealtimeResolver +# --------------------------------------------------------------------------- + + +class TestRealtimeResolver: + def test_returns_requested_asof_unchanged(self) -> None: + resolver = RealtimeResolver() + asof = date(2024, 6, 15) + result = resolver.resolve( + series_key="GDP", + obs_date=date(2024, 3, 31), + requested_asof=asof, + has_pit=True, + ) + assert result == asof + + def test_returns_requested_asof_for_non_pit(self) -> None: + resolver = RealtimeResolver() + asof = date(2024, 6, 15) + result = resolver.resolve( + series_key="SP500", + obs_date=date(2024, 3, 31), + requested_asof=asof, + has_pit=False, + ) + assert result == asof + + def test_view_property(self) -> None: + resolver = RealtimeResolver() + assert resolver.view == VintageView.realtime() + + def test_satisfies_protocol(self) -> None: + assert isinstance(RealtimeResolver(), VintageResolver) + + +# --------------------------------------------------------------------------- +# LatestResolver +# --------------------------------------------------------------------------- + + +class TestLatestResolver: + def test_returns_sentinel_for_pit_series(self) -> None: + resolver = LatestResolver() + result = resolver.resolve( + series_key="GDP", + obs_date=date(2024, 3, 31), + requested_asof=date(2024, 6, 15), + has_pit=True, + ) + assert result == date(2099, 12, 31) + + def test_returns_requested_asof_for_non_pit(self) -> None: + resolver = LatestResolver() + asof = date(2024, 6, 15) + result = resolver.resolve( + series_key="SP500", + obs_date=date(2024, 3, 31), + requested_asof=asof, + has_pit=False, + ) + assert result == asof + + def test_view_property(self) -> None: + resolver = LatestResolver() + assert resolver.view == VintageView.latest() + + def test_satisfies_protocol(self) -> None: + assert isinstance(LatestResolver(), VintageResolver) + + +# --------------------------------------------------------------------------- +# FrozenResolver +# --------------------------------------------------------------------------- + + +@pytest.fixture() +def revision_map() -> dict[tuple[str, date], list[date]]: + """Sample revision map: GDP Q1 2024 was released on three vintages.""" + return { + ("GDP", date(2024, 3, 31)): [ + date(2024, 4, 25), # advance + date(2024, 5, 30), # second + date(2024, 6, 27), # third + ], + ("GDP", date(2023, 12, 31)): [ + date(2024, 1, 25), # advance + date(2024, 2, 28), # second + ], + ("PAYEMS", date(2024, 3, 31)): [ + date(2024, 4, 5), # initial + date(2024, 5, 3), # revised + date(2024, 6, 7), # revised again + date(2024, 7, 5), # revised again + ], + } + + +class TestFrozenResolver: + def test_returns_nth_vintage(self, revision_map: dict) -> None: + resolver = FrozenResolver(revision_map=revision_map, n_releases=3) + result = resolver.resolve( + series_key="GDP", + obs_date=date(2024, 3, 31), + requested_asof=date(2024, 12, 1), + has_pit=True, + ) + assert result == date(2024, 6, 27) # 3rd release + + def test_returns_2nd_release(self, revision_map: dict) -> None: + resolver = FrozenResolver(revision_map=revision_map, n_releases=2) + result = resolver.resolve( + series_key="GDP", + obs_date=date(2024, 3, 31), + requested_asof=date(2024, 12, 1), + has_pit=True, + ) + assert result == date(2024, 5, 30) # 2nd release + + def test_returns_1st_release(self, revision_map: dict) -> None: + resolver = FrozenResolver(revision_map=revision_map, n_releases=1) + result = resolver.resolve( + series_key="GDP", + obs_date=date(2024, 3, 31), + requested_asof=date(2024, 12, 1), + has_pit=True, + ) + assert result == date(2024, 4, 25) # advance + + def test_falls_back_to_latest_when_fewer_releases(self, revision_map: dict) -> None: + resolver = FrozenResolver(revision_map=revision_map, n_releases=3) + result = resolver.resolve( + series_key="GDP", + obs_date=date(2023, 12, 31), + requested_asof=date(2024, 12, 1), + has_pit=True, + ) + # Only 2 releases exist, so returns the 2nd (latest available) + assert result == date(2024, 2, 28) + + def test_returns_requested_asof_for_non_pit(self, revision_map: dict) -> None: + resolver = FrozenResolver(revision_map=revision_map, n_releases=3) + asof = date(2024, 6, 15) + result = resolver.resolve( + series_key="SP500", + obs_date=date(2024, 3, 31), + requested_asof=asof, + has_pit=False, + ) + assert result == asof + + def test_returns_requested_asof_for_unknown_series(self, revision_map: dict) -> None: + resolver = FrozenResolver(revision_map=revision_map, n_releases=3) + asof = date(2024, 6, 15) + result = resolver.resolve( + series_key="UNKNOWN", + obs_date=date(2024, 3, 31), + requested_asof=asof, + has_pit=True, + ) + # Not in revision_map → fall through to requested_asof + assert result == asof + + def test_view_property(self, revision_map: dict) -> None: + resolver = FrozenResolver(revision_map=revision_map, n_releases=5) + assert resolver.view == VintageView.frozen(n=5) + + def test_satisfies_protocol(self, revision_map: dict) -> None: + assert isinstance( + FrozenResolver(revision_map=revision_map, n_releases=3), + VintageResolver, + ) + + def test_different_series_different_vintages(self, revision_map: dict) -> None: + resolver = FrozenResolver(revision_map=revision_map, n_releases=3) + gdp_result = resolver.resolve( + series_key="GDP", + obs_date=date(2024, 3, 31), + requested_asof=date(2024, 12, 1), + has_pit=True, + ) + payems_result = resolver.resolve( + series_key="PAYEMS", + obs_date=date(2024, 3, 31), + requested_asof=date(2024, 12, 1), + has_pit=True, + ) + assert gdp_result == date(2024, 6, 27) + assert payems_result == date(2024, 6, 7) + assert gdp_result != payems_result + + def test_empty_revision_map(self) -> None: + resolver = FrozenResolver(revision_map={}, n_releases=3) + asof = date(2024, 6, 15) + result = resolver.resolve( + series_key="GDP", + obs_date=date(2024, 3, 31), + requested_asof=asof, + has_pit=True, + ) + assert result == asof