import logging import os from collections.abc import Generator from contextlib import contextmanager from typing import Any import mlflow import mlflow.sklearn import pandas as pd from model_manager.sientia.exceptions import SientiaMlException class ModelServing: """ MLflow model serving wrapper. Thread-safety note: This class modifies global state (MLflow tracking URI and environment variables) during initialization. In multi-threaded environments, ensure that: 1. Instances are created with the same tracking_uri/credentials, OR 2. Instance creation is synchronized (e.g., using a lock), OR 3. Create a single instance and share it across threads The MLflow operations themselves (log_param, log_metric, etc.) are thread-safe when operating on different runs. """ def __init__( self, tracking_uri: str, username: str | None = None, password: str | None = None, ): """ Initialize ModelServing client. WARNING: This modifies global state (MLflow config and environment variables). Not thread-safe during initialization if different credentials are used. Args: tracking_uri: MLflow tracking server URI username: Optional MLflow username password: Optional MLflow password logger: Optional logger (currently unused) """ # Set tracking URI (modifies global MLflow state) mlflow.set_tracking_uri(tracking_uri) # Set credentials in environment variables (global state) if username is not None: os.environ['MLFLOW_TRACKING_USERNAME'] = username if password is not None: os.environ['MLFLOW_TRACKING_PASSWORD'] = password # Function to list runs for a given experiment def search_runs_by_name( self, experiment_names: list[str], order_by: None | list[str] = None ) -> pd.DataFrame | list: """ List runs for a specified MLflow experiment. Args: experiment_names (list[str]): List with experiment_names to retrieve runs from. Returns: Union[pd.DataFrame, list]: A DataFrame or list containing run information. Raises: SientiaMlException: If unable to search runs. """ try: runs = mlflow.search_runs(experiment_names=experiment_names, order_by=order_by) except SientiaMlException as e: logging.error(e) raise SientiaMlException(str(e)) from e return runs def set_experiment(self, experiment_identifier: str) -> None: """ Set the given experiment as the active experiment. Args: experiment_identifier (str): name or id of the experiment to be setted Raises: Exception: If setting the experiment fails. """ mlflow.set_experiment(experiment_identifier) def log_model(self, sk_model: Any, artifact_path: Any, **kwargs) -> None: """ Log a sklearn model. Args: sk_model: scikit-learn model to be saved. artifact_path: Run-relative artifact path. Returns: None Security Warning: The GitHub token is hardcoded. Consider moving to environment variable or using a secure secret management solution (e.g., K8s secrets). Raises: Exception: If logging the model fails. """ mlflow.sklearn.log_model( sk_model, artifact_path, extra_pip_requirements=[os.getenv('EXTRA_PIP_REQUIREMENTS')], **kwargs, ) def log_param(self, key: str, value: Any) -> None: """ Log a param in the active run. Args: key (str): Param name value (any): Param value Returns: None Raises: Exception: If logging the parameter fails. """ mlflow.log_param(key, value) def log_metric(self, key: str, value: Any) -> None: """ Log a metric in the active run. Args: key (str): Metric name value (any): Metric value Returns: None Raises: Exception: If logging the metric fails. """ mlflow.log_metric(key, value) def log_artifact( self, local_path: str, artifact_path: str | None = None, run_id: str | None = None ) -> None: """ Log an artifact. Args: local_path: Local path of the artifact to log. artifact_path: If provided, the directory in artifact_uri to write to. run_id: optional id of current run Returns: None Raises: Exception: If logging the artifact fails. """ mlflow.log_artifact(local_path=local_path, artifact_path=artifact_path, run_id=run_id) @contextmanager def save_experiment( self, run_id: str | None = None, experiment_id: str | None = None, run_name: str | None = None, nested: bool = False, tags: dict[str, Any] | None = None, description: str | None = None, log_system_metrics: bool | None = None, ) -> Generator[mlflow.ActiveRun, None, None]: """ Context manager to save an experiment, ensuring the run is properly closed. This prevents memory leaks by guaranteeing that MLflow runs are always ended, even if an exception occurs. Thread-safe when used with proper MLflow configuration. Args: run_id: If specified, get the run with the specified UUID and log parameters and metrics under that run. experiment_id: ID of the experiment under which to create the current run (applicable only when run_id is not specified). run_name: Name of new run. Used only when run_id is unspecified. nested: Controls whether run is nested in parent run. True creates a nested run. tags: An optional dictionary of string keys and values to set as tags on the run. If a run is being resumed, these tags are set on the resumed run. If a new run is being created, these tags are set on the new run. description: An optional string that populates the description box of the run. log_system_metrics: If True, system metrics will be logged. If None, we will check environment variable Yields: ActiveRun: object that acts as a context manager wrapping the run's state. Raises: Exception: If starting or ending the MLflow run fails. """ run = mlflow.start_run( run_id=run_id, experiment_id=experiment_id, run_name=run_name, nested=nested, tags=tags, description=description, log_system_metrics=log_system_metrics, ) try: yield run finally: # Ensure run is always ended, preventing resource leaks mlflow.end_run()