"""Prediction model adapters used through a small dependency-injection boundary.""" from __future__ import annotations from dataclasses import dataclass from pathlib import Path from typing import Protocol import pandas as pd from .artifact import ArtifactMetadata, load_model_artifact from .errors import InferenceError from .inference import DiagnosticModels, predict_test, validate_submission @dataclass(frozen=True) class PredictionBundle: submission: pd.DataFrame diagnostics: pd.DataFrame class PredictionModel(Protocol): def predict(self, frame: pd.DataFrame) -> PredictionBundle: ... class SklearnPredictionModel: """Runtime adapter around a verified, pre-trained model artifact.""" def __init__( self, models: DiagnosticModels, metadata: ArtifactMetadata | None = None ) -> None: self._models = models self.metadata = metadata @classmethod def from_artifact( cls, directory: Path, *, expected_model_version: str | None = None, ) -> "SklearnPredictionModel": try: artifact = load_model_artifact( directory, expected_model_version=expected_model_version ) except Exception as exc: raise InferenceError( "Nie udało się załadować zweryfikowanego artefaktu modelu.", hint="Sprawdź manifest, checksumę oraz zgodność wersji zależności.", ) from exc return cls(artifact.models, artifact.metadata) def predict(self, frame: pd.DataFrame) -> PredictionBundle: try: submission, diagnostics = predict_test(self._models, frame) validate_submission(submission, frame) except Exception as exc: raise InferenceError( "Model odrzucił pomiary podczas predykcji.", hint="Zweryfikuj kompletność silników oraz brak nietypowych wartości w widmie.", ) from exc diagnostics = diagnostics.copy() diagnostics["model_version"] = ( self.metadata.model_version if self.metadata else "unversioned" ) diagnostics["model_artifact_sha256"] = ( self.metadata.model_sha256 if self.metadata else "unavailable" ) return PredictionBundle(submission, diagnostics) class PrecomputedPredictionModel: """Read-only demo fallback; accepts only the exact precomputed key set.""" def __init__(self, submission: pd.DataFrame, diagnostics: pd.DataFrame) -> None: self._submission = submission.reset_index(drop=True).copy() self._diagnostics = diagnostics.reset_index(drop=True).copy() def predict(self, frame: pd.DataFrame) -> PredictionBundle: keys = ["engine_id", "cylinder"] if not frame[keys].reset_index(drop=True).equals(self._submission[keys]): raise InferenceError( "Tryb awaryjny obsługuje wyłącznie dołączony zestaw demonstracyjny.", hint="Przywróć artefakt modelu i uruchom ponownie aplikację, aby diagnozować własne pliki.", ) diagnostics = self._diagnostics.copy() diagnostics["model_version"] = "demo-precomputed" diagnostics["model_artifact_sha256"] = "unavailable" return PredictionBundle(self._submission.copy(), diagnostics)