Source code for mlflow.genai.scorers.registry

"""
Registered scorer functionality for MLflow GenAI.

This module provides functions to manage registered scorers that automatically
evaluate traces in MLflow experiments.
"""

import json
import warnings
from abc import ABCMeta, abstractmethod
from base64 import urlsafe_b64encode
from collections.abc import Callable
from functools import partial
from typing import TYPE_CHECKING, Any, NoReturn, Optional, TypeVar, cast
from urllib.parse import quote

from pydantic import BaseModel, ConfigDict, ValidationError
from typing_extensions import Self

from mlflow.entities import LifecycleStage
from mlflow.exceptions import MlflowException, RestException
from mlflow.genai.scorers.base import (
    SCORER_BACKEND_DATABRICKS,
    SCORER_BACKEND_TRACKING,
    Scorer,
    ScorerSamplingConfig,
)
from mlflow.protos.databricks_pb2 import (
    ALREADY_EXISTS,
    INTERNAL_ERROR,
    NOT_FOUND,
    RESOURCE_DOES_NOT_EXIST,
    ErrorCode,
)
from mlflow.tracking._tracking_service.utils import _get_store
from mlflow.tracking.fluent import _get_experiment_id
from mlflow.utils.databricks_utils import get_databricks_host_creds
from mlflow.utils.plugins import get_entry_points
from mlflow.utils.rest_utils import http_request, verify_rest_response
from mlflow.utils.uri import get_uri_scheme

if TYPE_CHECKING:
    from mlflow.genai.scorers.online.entities import OnlineScoringConfig

_T = TypeVar("_T")


class _DatabricksManagedScorerPayload(BaseModel):
    """Shared wire fields and parsing for Databricks managed-evals scorer payloads."""

    model_config = ConfigDict(extra="allow")

    serialized_scorer: str
    builtin: dict[str, str] | None = None
    custom: dict[str, Any] | None = None

    @classmethod
    def _parse(cls, value: Any, response_field: str) -> Self:
        try:
            return cls.model_validate(value)
        except ValidationError as e:
            raise MlflowException(
                f"Failed to parse managed scorer response field `{response_field}`: {e}",
                INTERNAL_ERROR,
            ) from e

    @classmethod
    def _parse_list(cls, value: Any, response_field: str) -> list[Self]:
        if not isinstance(value, list):
            raise MlflowException(
                f"Failed to parse managed scorer response field `{response_field}`: "
                "expected a list.",
                INTERNAL_ERROR,
            )
        return [cls._parse(item, response_field) for item in value]


class _DatabricksScheduledScorerConfig(_DatabricksManagedScorerPayload):
    """Mutable scorer configuration returned by managed scheduled-scorer endpoints."""

    name: str
    sample_rate: float | None = None
    filter_string: str | None = None
    scorer_version: int | None = None

    @classmethod
    def from_scorer(
        cls,
        *,
        name: str,
        scorer: Scorer,
        sample_rate: float | None,
        filter_string: str | None,
    ) -> Self:
        serialized_scorer = scorer.model_dump()
        # Preserve registration-time validation that the serialized scorer can be reconstructed.
        Scorer.model_validate(serialized_scorer)
        config: dict[str, Any] = {
            "name": name,
            "serialized_scorer": json.dumps(serialized_scorer),
        }
        if serialized_scorer.get("builtin_scorer_class"):
            config["builtin"] = {"name": name}
        else:
            config["custom"] = {}
        if sample_rate is not None:
            config["sample_rate"] = sample_rate
        if filter_string is not None:
            config["filter_string"] = filter_string
        return cls.model_validate(config)

    @classmethod
    def from_list_response(cls, response: dict[str, Any]) -> list[Self]:
        scheduled_scorers = response.get("scheduled_scorers", {})
        if not isinstance(scheduled_scorers, dict):
            raise MlflowException(
                "Failed to parse managed scorer response field `scheduled_scorers`: "
                "expected an object.",
                INTERNAL_ERROR,
            )
        return cls._parse_list(
            scheduled_scorers.get("scorers", []),
            "scheduled_scorers.scorers",
        )


class _DatabricksScorerVersion(_DatabricksManagedScorerPayload):
    """Immutable scorer definition returned by managed scorer-version endpoints."""

    name: str
    display_name: str | None = None
    scorer_version: int
    create_time: str | None = None

    @classmethod
    def from_response(cls, response: dict[str, Any]) -> Self:
        return cls._parse(response, "scorer_version")

    @classmethod
    def from_list_response(cls, response: dict[str, Any]) -> list[Self]:
        return cls._parse_list(response.get("scorer_versions", []), "scorer_versions")


class UnsupportedScorerStoreURIException(MlflowException):
    """Exception thrown when building a scorer store with an unsupported URI"""

    def __init__(self, unsupported_uri, supported_uri_schemes):
        message = (
            f"Scorer registration functionality is unavailable; got unsupported URI"
            f" '{unsupported_uri}' for scorer data storage. Supported URI schemes are:"
            f" {supported_uri_schemes}."
        )
        super().__init__(message)
        self.supported_uri_schemes = supported_uri_schemes


class AbstractScorerStore(metaclass=ABCMeta):
    """
    Abstract class defining the interface for scorer store implementations.

    This class defines the API interface for scorer operations that can be implemented
    by different backend stores (e.g., MLflow tracking store, Databricks API).
    """

    @abstractmethod
    def register_scorer(self, experiment_id: str | None, scorer: Scorer) -> int | None:
        """
        Register a scorer for an experiment.

        Args:
            experiment_id: The ID of the Experiment containing the scorer.
            scorer: The scorer object.

        Returns:
            The registered scorer version. If versioning is not supported, return None.
        """

    @abstractmethod
    def list_scorers(self, experiment_id) -> list["Scorer"]:
        """
        List all scorers for an experiment.

        Args:
            experiment_id: The ID of the Experiment containing the scorer.

        Returns:
            List of mlflow.genai.scorers.Scorer objects (latest version for each scorer name).
        """

    @abstractmethod
    def get_scorer(self, experiment_id, name, version=None) -> "Scorer":
        """
        Get a specific scorer for an experiment.

        Args:
            experiment_id: The ID of the Experiment containing the scorer.
            name: The scorer name.
            version: The scorer version. If None, returns the scorer with maximum version.

        Returns:
            A list of tuple, each tuple contains `mlflow.genai.scorers.Scorer` object.

        Raises:
            mlflow.MlflowException: If scorer is not found.
        """

    @abstractmethod
    def list_scorer_versions(self, experiment_id, name) -> list[tuple["Scorer", int]]:
        """
        List all versions of a specific scorer for an experiment.

        Args:
            experiment_id: The ID of the Experiment containing the scorer.
            name: The scorer name.

        Returns:
            A list of tuple, each tuple contains `mlflow.genai.scorers.Scorer` object
            and the version number.

        Raises:
            mlflow.MlflowException: If scorer is not found.
        """

    @abstractmethod
    def delete_scorer(self, experiment_id, name, version):
        """
        Delete a scorer by name and optional version.

        Args:
            experiment_id: The ID of the Experiment containing the scorer.
            name: The scorer name.
            version: The scorer version to delete.

        Raises:
            mlflow.MlflowException: If scorer is not found.
        """


class ScorerStoreRegistry:
    """
    Scheme-based registry for scorer store implementations.

    This class allows the registration of a function or class to provide an
    implementation for a given scheme of `store_uri` through the `register`
    methods. Implementations declared though the entrypoints
    `mlflow.scorer_store` group can be automatically registered through the
    `register_entrypoints` method.

    When instantiating a store through the `get_store` method, the scheme of
    the store URI provided (or inferred from environment) will be used to
    select which implementation to instantiate, which will be called with same
    arguments passed to the `get_store` method.
    """

    def __init__(self):
        self._registry = {}
        self.group_name = "mlflow.scorer_store"

    def register(self, scheme, store_builder):
        self._registry[scheme] = store_builder

    def register_entrypoints(self):
        """Register scorer stores provided by other packages"""
        for entrypoint in get_entry_points(self.group_name):
            try:
                self.register(entrypoint.name, entrypoint.load())
            except (AttributeError, ImportError) as exc:
                warnings.warn(
                    'Failure attempting to register scorer store for scheme "{}": {}'.format(
                        entrypoint.name, str(exc)
                    ),
                    stacklevel=2,
                )

    def get_store_builder(self, store_uri):
        """Get a store from the registry based on the scheme of store_uri

        Args:
            store_uri: The store URI. If None, it will be inferred from the environment. This
                URI is used to select which scorer store implementation to instantiate
                and is passed to the constructor of the implementation.

        Returns:
            A function that returns an instance of
            ``mlflow.genai.scorers.registry.AbstractScorerStore`` that fulfills the store
            URI requirements.
        """
        scheme = store_uri if store_uri == "databricks" else get_uri_scheme(store_uri)
        try:
            store_builder = self._registry[scheme]
        except KeyError:
            raise UnsupportedScorerStoreURIException(
                unsupported_uri=store_uri, supported_uri_schemes=list(self._registry.keys())
            )
        return store_builder

    def get_store(self, tracking_uri=None):
        from mlflow.tracking._tracking_service import utils

        resolved_store_uri = utils._resolve_tracking_uri(tracking_uri)
        builder = self.get_store_builder(resolved_store_uri)
        return builder(tracking_uri=resolved_store_uri)


class MlflowTrackingStore(AbstractScorerStore):
    """
    MLflow tracking store that provides scorer functionality through the tracking store.
    This store delegates all scorer operations to the underlying tracking store.
    """

    def __init__(self, tracking_uri=None):
        self._tracking_store = _get_store(tracking_uri)

    def register_scorer(self, experiment_id: str | None, scorer: Scorer) -> int | None:
        serialized_scorer = json.dumps(scorer.model_dump())
        experiment_id = experiment_id or _get_experiment_id()
        version = self._tracking_store.register_scorer(
            experiment_id, scorer.name, serialized_scorer
        )
        self._hydrate_scorer(scorer, experiment_id, online_config=None)
        return version

    def _hydrate_scorer(
        self,
        scorer: Scorer,
        experiment_id: str,
        online_config: Optional["OnlineScoringConfig"] = None,
    ) -> None:
        """
        Hydrate a scorer with runtime state from the tracking store.

        Args:
            scorer: The scorer to hydrate.
            experiment_id: The experiment ID the scorer belongs to.
            online_config: Optional OnlineScoringConfig from the tracking store.
        """
        scorer._registered_backend = SCORER_BACKEND_TRACKING
        scorer._experiment_id = experiment_id
        if online_config is not None:
            scorer._sampling_config = ScorerSamplingConfig(
                sample_rate=online_config.sample_rate,
                filter_string=online_config.filter_string,
            )

    def list_scorers(self, experiment_id) -> list["Scorer"]:
        from mlflow.genai.scorers import Scorer

        experiment_id = experiment_id or _get_experiment_id()
        scorer_versions = self._tracking_store.list_scorers(experiment_id)
        scorer_ids = [sv.scorer_id for sv in scorer_versions]
        online_configs_list = (
            self._tracking_store.get_online_scoring_configs(scorer_ids) if scorer_ids else []
        )
        # Each scorer has at most one online configuration, guaranteed by the server
        online_configs = {c.scorer_id: c for c in online_configs_list}
        scorers = []
        for scorer_version in scorer_versions:
            scorer = Scorer.model_validate(scorer_version.serialized_scorer)
            online_config = online_configs.get(scorer_version.scorer_id)
            self._hydrate_scorer(scorer, experiment_id, online_config)
            scorers.append(scorer)
        return scorers

    def get_scorer(self, experiment_id, name, version=None) -> "Scorer":
        from mlflow.genai.scorers import Scorer

        experiment_id = experiment_id or _get_experiment_id()
        scorer_version = self._tracking_store.get_scorer(experiment_id, name, version)
        online_configs_list = self._tracking_store.get_online_scoring_configs([
            scorer_version.scorer_id
        ])
        # Each scorer has at most one online configuration, guaranteed by the server
        online_config = online_configs_list[0] if online_configs_list else None
        scorer = Scorer.model_validate(scorer_version.serialized_scorer)
        self._hydrate_scorer(scorer, experiment_id, online_config)
        return scorer

    def list_scorer_versions(self, experiment_id, name) -> list[tuple[Scorer, int]]:
        from mlflow.genai.scorers import Scorer

        experiment_id = experiment_id or _get_experiment_id()
        scorer_versions = self._tracking_store.list_scorer_versions(experiment_id, name)
        scorer_ids = list({sv.scorer_id for sv in scorer_versions})
        online_configs_list = (
            self._tracking_store.get_online_scoring_configs(scorer_ids) if scorer_ids else []
        )
        # Each scorer has at most one online configuration, guaranteed by the server
        online_configs = {c.scorer_id: c for c in online_configs_list}
        scorers = []
        for scorer_version in scorer_versions:
            scorer = Scorer.model_validate(scorer_version.serialized_scorer)
            online_config = online_configs.get(scorer_version.scorer_id)
            self._hydrate_scorer(scorer, experiment_id, online_config)
            scorers.append((scorer, scorer_version.scorer_version))
        return scorers

    def delete_scorer(self, experiment_id, name, version):
        if version is None:
            raise MlflowException.invalid_parameter_value(
                "You must set `version` argument to either an integer or 'all'."
            )
        if version == "all":
            version = None

        experiment_id = experiment_id or _get_experiment_id()
        return self._tracking_store.delete_scorer(experiment_id, name, version)

    def upsert_online_scoring_config(
        self,
        *,
        scorer: Scorer,
        experiment_id: str,
        sample_rate: float,
        filter_string: str | None = None,
    ) -> Scorer:
        """
        Create or update the online scoring configuration for a registered scorer.

        Args:
            scorer: The scorer instance to update.
            experiment_id: The ID of the MLflow experiment containing the scorer.
            sample_rate: The sampling rate (0.0 to 1.0).
            filter_string: Optional filter string.

        Returns:
            A copy of the scorer with updated sampling configuration.

        Raises:
            MlflowException: If the scorer is not registered.
        """
        if scorer._registered_backend is None:
            raise MlflowException.invalid_parameter_value(
                "Cannot start/update a scorer that is not registered. "
                "Please call register() first before calling start()/update(), "
                "or use get_scorer() to load a registered scorer."
            )

        self._tracking_store.upsert_online_scoring_config(
            experiment_id=experiment_id,
            scorer_name=scorer.name,
            sample_rate=sample_rate,
            filter_string=filter_string,
        )

        return self.get_scorer(experiment_id, scorer.name)


class DatabricksStore(AbstractScorerStore):
    """
    Databricks store that provides scorer functionality through the Databricks API.
    This store delegates current scorer operations to the Databricks agents API and uses the
    managed-evals API for versioned operations.
    """

    # TODO: Extract managed-evals request and pagination handling into a shared
    # ManagedEvalsClient used by other managed-evals integrations.
    _MANAGED_EVALS_BASE = "/api/2.0/managed-evals"
    _MANAGED_EVALS_SCHEDULED_SCORERS_BASE = f"{_MANAGED_EVALS_BASE}/scheduled-scorers"

    def __init__(self, tracking_uri=None):
        self._tracking_uri = tracking_uri
        self.get_host_creds = partial(get_databricks_host_creds, tracking_uri)

    @staticmethod
    def _resolve_experiment_id(experiment_id: str | None) -> str:
        if resolved_experiment_id := experiment_id or _get_experiment_id():
            return resolved_experiment_id
        raise MlflowException(
            "No active experiment found. Set an experiment using `mlflow.set_experiment`, "
            "or pass `experiment_id`."
        )

    @staticmethod
    def _scorer_not_found(
        experiment_id: str, name: str, version: int | None = None
    ) -> MlflowException:
        version_description = f" and version {version}" if version is not None else ""
        return MlflowException(
            f"Scorer with name '{name}'{version_description} not found for experiment "
            f"{experiment_id}.",
            RESOURCE_DOES_NOT_EXIST,
        )

    @classmethod
    def _raise_scorer_rest_error(
        cls,
        error: RestException,
        *,
        experiment_id: str,
        name: str,
        version: int | None = None,
    ) -> NoReturn:
        message = str(error.json.get("message", ""))
        if (
            error.error_code != ErrorCode.Name(NOT_FOUND)
            or "versioning is not enabled" in message.lower()
        ):
            raise error
        missing_version = version if "scorer version " in message.lower() else None
        raise cls._scorer_not_found(experiment_id, name, missing_version) from error

    def _validate_experiment_is_active(self, experiment_id: str) -> None:
        experiment = _get_store(self._tracking_uri).get_experiment(experiment_id)
        if experiment.lifecycle_stage != LifecycleStage.ACTIVE:
            raise MlflowException.invalid_parameter_value(
                f"The experiment {experiment.experiment_id} must be in the 'active' state. "
                f"Current state is {experiment.lifecycle_stage}."
            )

    @staticmethod
    def _encode_path_param(value: str) -> str:
        return quote(str(value), safe="")

    @staticmethod
    def _scorer_resource_key(name: str) -> str:
        return urlsafe_b64encode(name.encode("utf-8")).decode("ascii").rstrip("=")

    def _scheduled_scorers_endpoint(self, experiment_id: str) -> str:
        return (
            f"{self._MANAGED_EVALS_SCHEDULED_SCORERS_BASE}/{self._encode_path_param(experiment_id)}"
        )

    @staticmethod
    def _validate_version(version: object) -> int:
        if isinstance(version, bool) or not isinstance(version, int) or version <= 0:
            raise MlflowException.invalid_parameter_value(
                f"`version` must be a positive integer, got {version!r}."
            )
        return version

    def _scorer_version_endpoint(self, experiment_id: str, name: str, version: int) -> str:
        version = self._validate_version(version)
        return f"{self._scorer_versions_endpoint(experiment_id, name)}/{version}"

    def _scorer_versions_endpoint(self, experiment_id: str, name: str) -> str:
        return (
            f"{self._MANAGED_EVALS_BASE}/experiments/{self._encode_path_param(experiment_id)}"
            f"/scorers/{self._scorer_resource_key(name)}/versions"
        )

    def _request(
        self,
        method: str,
        endpoint: str,
        *,
        json_body: dict[str, Any] | None = None,
        params: dict[str, Any] | None = None,
    ) -> dict[str, Any]:
        response = http_request(
            host_creds=self.get_host_creds(),
            endpoint=endpoint,
            method=method,
            json=json_body,
            params=params,
        )
        verify_rest_response(response, endpoint)
        if not response.text:
            return {}
        return cast(dict[str, Any], response.json())

    def _get_paginated_results(
        self,
        endpoint: str,
        extract_items: Callable[[dict[str, Any]], list[_T]],
    ) -> list[_T]:
        items: list[_T] = []
        page_token: str | None = None
        while True:
            params = {"page_token": page_token} if page_token else None
            response = self._request("GET", endpoint, params=params)
            items.extend(extract_items(response))

            next_page_token = response.get("next_page_token")
            if next_page_token is None:
                return items
            if (
                not isinstance(next_page_token, str)
                or not next_page_token
                or next_page_token == page_token
            ):
                raise MlflowException(
                    "Paginated response contained an invalid `next_page_token`.",
                    INTERNAL_ERROR,
                )
            page_token = next_page_token

    def _list_current_scorer_configs(
        self, experiment_id: str
    ) -> list[_DatabricksScheduledScorerConfig]:
        return self._get_paginated_results(
            self._scheduled_scorers_endpoint(experiment_id),
            _DatabricksScheduledScorerConfig.from_list_response,
        )

    def _patch_current_scorer_configs(
        self, experiment_id: str, configs: list[_DatabricksScheduledScorerConfig]
    ) -> list[_DatabricksScheduledScorerConfig]:
        response = self._request(
            "PATCH",
            self._scheduled_scorers_endpoint(experiment_id),
            json_body={
                "scheduled_scorers": {
                    "scorers": [config.model_dump(exclude_unset=True) for config in configs]
                },
                "update_mask": "scheduled_scorers.scorers",
            },
        )
        return _DatabricksScheduledScorerConfig.from_list_response(response)

    def _create_current_scorer_configs(
        self, experiment_id: str, configs: list[_DatabricksScheduledScorerConfig]
    ) -> list[_DatabricksScheduledScorerConfig]:
        response = self._request(
            "POST",
            self._scheduled_scorers_endpoint(experiment_id),
            json_body={
                "scheduled_scorers": {
                    "scorers": [config.model_dump(exclude_unset=True) for config in configs]
                }
            },
        )
        return _DatabricksScheduledScorerConfig.from_list_response(response)

    def _upsert_registered_scorer_config(
        self,
        experiment_id: str,
        registered_config: _DatabricksScheduledScorerConfig,
    ) -> list[_DatabricksScheduledScorerConfig]:
        configs = self._merge_registered_scorer_config(
            self._list_current_scorer_configs(experiment_id), registered_config
        )
        try:
            return self._patch_current_scorer_configs(experiment_id, configs)
        except RestException as e:
            if e.error_code != ErrorCode.Name(NOT_FOUND):
                raise
        try:
            return self._create_current_scorer_configs(experiment_id, configs)
        except RestException as e:
            if e.error_code != ErrorCode.Name(ALREADY_EXISTS):
                raise

        configs = self._merge_registered_scorer_config(
            self._list_current_scorer_configs(experiment_id), registered_config
        )
        return self._patch_current_scorer_configs(experiment_id, configs)

    def _find_current_scorer_config(
        self, experiment_id: str, name: str
    ) -> _DatabricksScheduledScorerConfig:
        configs = self._list_current_scorer_configs(experiment_id)
        if (index := self._find_scorer_config_index(configs, name)) is not None:
            return configs[index]
        raise MlflowException(
            f"Scorer with name '{name}' not found for experiment {experiment_id}.",
            RESOURCE_DOES_NOT_EXIST,
        )

    @staticmethod
    def _find_scorer_config_index(
        configs: list[_DatabricksScheduledScorerConfig], name: str
    ) -> int | None:
        return next((i for i, config in enumerate(configs) if config.name == name), None)

    @classmethod
    def _merge_registered_scorer_config(
        cls,
        configs: list[_DatabricksScheduledScorerConfig],
        registered_config: _DatabricksScheduledScorerConfig,
    ) -> list[_DatabricksScheduledScorerConfig]:
        configs = list(configs)
        index = cls._find_scorer_config_index(configs, registered_config.name)
        if index is None:
            configs.append(registered_config)
            return configs

        current_payload = configs[index].model_dump(exclude_unset=True)
        registered_payload = registered_config.model_dump(exclude_unset=True)
        for field in ("sample_rate", "filter_string"):
            if field in current_payload:
                registered_payload[field] = current_payload[field]
            else:
                registered_payload.pop(field, None)
        if configs[index].scorer_version is not None:
            registered_payload["scorer_version"] = configs[index].scorer_version
        else:
            registered_payload.pop("scorer_version", None)
        configs[index] = _DatabricksScheduledScorerConfig.model_validate(registered_payload)
        return configs

    # Private functions for internal use by Scorer methods
    @staticmethod
    def list_scheduled_scorers(experiment_id):
        try:
            from databricks.agents.scorers import list_scheduled_scorers
        except ImportError as e:
            raise ImportError(_ERROR_MSG) from e

        return list_scheduled_scorers(experiment_id=experiment_id)

    @staticmethod
    def get_scheduled_scorer(name, experiment_id):
        try:
            from databricks.agents.scorers import get_scheduled_scorer
        except ImportError as e:
            raise ImportError(_ERROR_MSG) from e

        return get_scheduled_scorer(
            scheduled_scorer_name=name,
            experiment_id=experiment_id,
        )

    @staticmethod
    def delete_scheduled_scorer(experiment_id, name):
        try:
            from databricks.agents.scorers import delete_scheduled_scorer
        except ImportError as e:
            raise ImportError(_ERROR_MSG) from e

        delete_scheduled_scorer(
            experiment_id=experiment_id,
            scheduled_scorer_name=name,
        )

    def update_registered_scorer(
        self,
        *,
        name: str,
        scorer: Scorer | None = None,
        sample_rate: float | None = None,
        filter_string: str | None = None,
        experiment_id: str | None = None,
    ) -> Scorer:
        """Update scheduling fields without changing the current scorer definition."""
        experiment_id = self._resolve_experiment_id(experiment_id)
        configs = self._list_current_scorer_configs(experiment_id)
        index = self._find_scorer_config_index(configs, name)
        if index is None:
            raise MlflowException(
                f"Scorer with name '{name}' not found for experiment {experiment_id}.",
                RESOURCE_DOES_NOT_EXIST,
            )
        updates = {}
        if sample_rate is not None:
            updates["sample_rate"] = sample_rate
        if filter_string is not None:
            updates["filter_string"] = filter_string
        configs[index] = configs[index].model_copy(update=updates)

        for config in self._patch_current_scorer_configs(experiment_id, configs):
            if config.name == name:
                return Scorer.model_validate_json(
                    config.serialized_scorer
                )._set_registration_metadata(
                    backend=SCORER_BACKEND_DATABRICKS,
                    experiment_id=experiment_id,
                    sampling_config=ScorerSamplingConfig(
                        sample_rate=config.sample_rate,
                        filter_string=config.filter_string,
                    ),
                )
        raise MlflowException(f"Updated scheduled scorer response did not include '{name}'.")

    def register_scorer(self, experiment_id: str | None, scorer: Scorer) -> int | None:
        experiment_id = self._resolve_experiment_id(experiment_id)
        registered_config = _DatabricksScheduledScorerConfig.from_scorer(
            name=scorer.name,
            scorer=scorer,
            sample_rate=0.0,
            filter_string=None,
        )
        response_configs = self._upsert_registered_scorer_config(experiment_id, registered_config)

        response_config = next(
            (config for config in response_configs if config.name == scorer.name), None
        )
        if response_config is None:
            raise MlflowException(f"Scheduled scorer response did not include '{scorer.name}'.")

        scorer._set_registration_metadata(
            backend=SCORER_BACKEND_DATABRICKS,
            experiment_id=experiment_id,
            sampling_config=ScorerSamplingConfig(
                sample_rate=response_config.sample_rate,
                filter_string=response_config.filter_string,
            ),
        )
        return response_config.scorer_version

    def list_scorers(self, experiment_id) -> list["Scorer"]:
        experiment_id = self._resolve_experiment_id(experiment_id)
        # Get scheduled scorers from the server
        scheduled_scorers = self.list_scheduled_scorers(experiment_id)
        self._validate_experiment_is_active(experiment_id)

        # Convert to Scorer instances with registration info
        return [
            scheduled_scorer.scorer._set_registration_metadata(
                backend=SCORER_BACKEND_DATABRICKS,
                experiment_id=experiment_id,
                sampling_config=ScorerSamplingConfig(
                    sample_rate=scheduled_scorer.sample_rate,
                    filter_string=scheduled_scorer.filter_string,
                ),
            )
            for scheduled_scorer in scheduled_scorers
        ]

    def get_scorer(self, experiment_id, name, version=None) -> "Scorer":
        if version is not None:
            experiment_id = self._resolve_experiment_id(experiment_id)
            try:
                response = self._request(
                    "GET", self._scorer_version_endpoint(experiment_id, name, version)
                )
            except RestException as e:
                self._raise_scorer_rest_error(
                    e,
                    experiment_id=experiment_id,
                    name=name,
                    version=version,
                )
            version_config = _DatabricksScorerVersion.from_response(response)
            current_config = self._find_current_scorer_config(experiment_id, name)
            return Scorer.model_validate_json(
                version_config.serialized_scorer
            )._set_registration_metadata(
                backend=SCORER_BACKEND_DATABRICKS,
                experiment_id=experiment_id,
                sampling_config=ScorerSamplingConfig(
                    sample_rate=current_config.sample_rate,
                    filter_string=current_config.filter_string,
                ),
            )

        # Get the scheduled scorer from the server
        experiment_id = self._resolve_experiment_id(experiment_id)
        try:
            scheduled_scorer = self.get_scheduled_scorer(name, experiment_id)
        except ValueError as e:
            if "No registered scorer found with name" not in str(e):
                raise
            raise self._scorer_not_found(experiment_id, name) from e

        # Extract the scorer and set registration fields
        return scheduled_scorer.scorer._set_registration_metadata(
            backend=SCORER_BACKEND_DATABRICKS,
            experiment_id=experiment_id,
            sampling_config=ScorerSamplingConfig(
                sample_rate=scheduled_scorer.sample_rate,
                filter_string=scheduled_scorer.filter_string,
            ),
        )

    def list_scorer_versions(self, experiment_id, name) -> list[tuple["Scorer", int]]:
        experiment_id = self._resolve_experiment_id(experiment_id)
        try:
            configs = self._get_paginated_results(
                self._scorer_versions_endpoint(experiment_id, name),
                _DatabricksScorerVersion.from_list_response,
            )
        except RestException as e:
            self._raise_scorer_rest_error(e, experiment_id=experiment_id, name=name)

        current_config = self._find_current_scorer_config(experiment_id, name)
        return [
            (
                Scorer.model_validate_json(config.serialized_scorer)._set_registration_metadata(
                    backend=SCORER_BACKEND_DATABRICKS,
                    experiment_id=experiment_id,
                    sampling_config=ScorerSamplingConfig(
                        sample_rate=current_config.sample_rate,
                        filter_string=current_config.filter_string,
                    ),
                ),
                config.scorer_version,
            )
            for config in configs
        ]

    def delete_scorer(self, experiment_id, name, version):
        if version is None or version == "all":
            experiment_id = self._resolve_experiment_id(experiment_id)
            try:
                return DatabricksStore.delete_scheduled_scorer(experiment_id, name)
            except ValueError as e:
                if "No registered scorer found with name" not in str(e):
                    raise
                raise self._scorer_not_found(experiment_id, name) from e

        experiment_id = self._resolve_experiment_id(experiment_id)
        try:
            self._request(
                "DELETE",
                self._scorer_version_endpoint(experiment_id, name, version),
            )
        except RestException as e:
            self._raise_scorer_rest_error(
                e,
                experiment_id=experiment_id,
                name=name,
                version=version,
            )


# Create the global scorer store registry instance
_scorer_store_registry = ScorerStoreRegistry()


def _register_scorer_stores():
    """Register the default scorer store implementations"""
    from mlflow.store.db.db_types import DATABASE_ENGINES

    # Register for database schemes (these will use MlflowTrackingStore)
    for scheme in DATABASE_ENGINES + ["http", "https"]:
        _scorer_store_registry.register(scheme, MlflowTrackingStore)

    # Register Databricks store
    _scorer_store_registry.register("databricks", DatabricksStore)

    # Register entrypoints for custom implementations
    _scorer_store_registry.register_entrypoints()


# Register the default stores
_register_scorer_stores()


def _get_scorer_store(tracking_uri=None):
    """Get a scorer store from the registry"""
    return _scorer_store_registry.get_store(tracking_uri)


_ERROR_MSG = (
    "The `databricks-agents` package is required to register scorers. "
    "Please install it with `pip install databricks-agents`."
)


[docs]def list_scorers(*, experiment_id: str | None = None) -> list[Scorer]: """ List all registered scorers for an experiment. This function retrieves all scorers that have been registered in the specified experiment. For each scorer name, only the latest version is returned. The function automatically determines the appropriate backend store (MLflow tracking store, Databricks, etc.) based on the current MLflow configuration and experiment location. Args: experiment_id (str, optional): The ID of the MLflow experiment containing the scorers. If None, uses the currently active experiment as determined by :func:`mlflow.get_experiment_by_name` or :func:`mlflow.set_experiment`. Returns: list[Scorer]: A list of Scorer objects, each representing the latest version of a registered scorer with its current configuration. The list may be empty if no scorers have been registered in the experiment. Raises: mlflow.MlflowException: If the experiment doesn't exist or if there are issues with the backend store connection. Example: .. code-block:: python from mlflow.genai.scorers import list_scorers # List all scorers in the current experiment scorers = list_scorers() # List all scorers in a specific experiment scorers = list_scorers(experiment_id="123") # Process the returned scorers for scorer in scorers: print(f"Scorer: {scorer.name}") Note: - Only the latest version of each scorer is returned. - This function works with both OSS MLflow tracking backend and Databricks backend. """ store = _get_scorer_store() return store.list_scorers(experiment_id)
def list_scorer_versions( *, name: str, experiment_id: str | None = None ) -> list[tuple[Scorer, int]]: """ List all versions of a specific scorer for an experiment. This function retrieves all versions of a scorer with the specified name from the given experiment. The function returns a list of tuples, where each tuple contains a Scorer instance and its corresponding version number. Args: name (str): The name of the scorer to list versions for. This must match exactly with the name used during scorer registration. experiment_id (str, optional): The ID of the MLflow experiment containing the scorer. If None, uses the currently active experiment as determined by :func:`mlflow.get_experiment_by_name` or :func:`mlflow.set_experiment`. Returns: list[tuple[Scorer, int]]: A list of tuples, where each tuple contains: - A Scorer object representing the scorer at that specific version - An integer representing the version number (1, 2, 3, etc.). The list may be empty if no versions of the scorer exist. Raises: mlflow.MlflowException: If the scorer with the specified name is not found in the experiment, if the experiment doesn't exist, or if there are issues with the backend store. """ store = _get_scorer_store() return store.list_scorer_versions(experiment_id, name)
[docs]def get_scorer( *, name: str, experiment_id: str | None = None, version: int | None = None ) -> Scorer: """ Retrieve a specific registered scorer by name and optional version. This function retrieves a single Scorer instance from the specified experiment. If no version is specified, it returns the latest (highest version number) scorer with the given name. Args: name (str): The name of the registered scorer to retrieve. This must match exactly with the name used during scorer registration. experiment_id (str, optional): The ID of the MLflow experiment containing the scorer. If None, uses the currently active experiment as determined by :func:`mlflow.get_experiment_by_name` or :func:`mlflow.set_experiment`. version (int, optional): The specific version of the scorer to retrieve. If None, returns the scorer with the highest version number (latest version). Returns: Scorer: A Scorer object representing the requested scorer. Raises: mlflow.MlflowException: If the scorer with the specified name is not found in the experiment, if the specified version doesn't exist, if the experiment doesn't exist, or if there are issues with the backend store connection. Example: .. code-block:: python from mlflow.genai.scorers import get_scorer # Get the latest version of a scorer latest_scorer = get_scorer(name="accuracy_scorer") # Get a specific version of a scorer v2_scorer = get_scorer(name="safety_scorer", version=2) # Get a scorer from a specific experiment scorer = get_scorer(name="relevance_scorer", experiment_id="123") Note: - When no version is specified, the function automatically returns the latest version - This function works with both OSS MLflow tracking backend and Databricks backend. """ store = _get_scorer_store() return store.get_scorer(experiment_id, name, version)
[docs]def delete_scorer( *, name: str, experiment_id: str | None = None, version: int | str | None = None, ) -> None: """ Delete a registered scorer from the MLflow experiment. This function permanently removes scorer registrations. The behavior of this function varies depending on the backend store and version parameter: **OSS MLflow Tracking Backend:** - Supports versioning with granular deletion options - Can delete specific versions or all versions of a scorer by setting `version` parameter to "all" **Databricks Backend:** - Supports deleting a specific version - Supports deleting all versions with `version="all"` - For backwards compatibility, `version=None` also deletes all versions Args: name (str): The name of the scorer to delete. This must match exactly with the name used during scorer registration. experiment_id (str, optional): The ID of the MLflow experiment containing the scorer. If None, uses the currently active experiment as determined by :func:`mlflow.get_experiment_by_name` or :func:`mlflow.set_experiment`. version (int | str | None, optional): The version(s) to delete. An integer deletes that specific version, and the string `"all"` deletes all versions. The OSS backend requires this argument. For backwards compatibility, the Databricks backend treats `None` as `"all"`. Raises: mlflow.MlflowException: If the scorer with the specified name is not found in the experiment, if the specified version doesn't exist, or if versioning is not supported for the current backend. Example: .. code-block:: python from mlflow.genai.scorers import delete_scorer # Delete a specific version of a scorer delete_scorer(name="safety_scorer", version=2) # Delete all versions of a scorer delete_scorer(name="relevance_scorer", version="all") # Delete a scorer from a specific experiment delete_scorer(name="harmfulness_scorer", experiment_id="123", version=1) """ store = _get_scorer_store() return store.delete_scorer(experiment_id, name, version)