import logging import threading import urllib import uuid from typing import Any import sqlalchemy from sqlalchemy import select from sqlalchemy.orm import Session from mlflow.entities.model_registry.model_version_stages import ( ALL_STAGES, DEFAULT_STAGES_FOR_GET_LATEST_VERSIONS, STAGE_ARCHIVED, STAGE_DELETED_INTERNAL, get_canonical_stage, ) from mlflow.entities.model_registry.prompt_version import IS_PROMPT_TAG_KEY from mlflow.entities.webhook import Webhook, WebhookEvent, WebhookStatus from mlflow.exceptions import MlflowException from mlflow.prompt.registry_utils import handle_resource_already_exist_error, has_prompt_tag from mlflow.protos.databricks_pb2 import ( INVALID_PARAMETER_VALUE, INVALID_STATE, RESOURCE_ALREADY_EXISTS, RESOURCE_DOES_NOT_EXIST, ) from mlflow.store.artifact.utils.models import _parse_model_uri from mlflow.store.db.utils import ( _all_tables_exist, _get_managed_session_maker, _get_routing_session_maker, _initialize_tables, create_sqlalchemy_engine_with_retry, ) from mlflow.store.entities.paged_list import PagedList from mlflow.store.model_registry import ( SEARCH_MODEL_VERSION_MAX_RESULTS_DEFAULT, SEARCH_MODEL_VERSION_MAX_RESULTS_THRESHOLD, SEARCH_REGISTERED_MODEL_MAX_RESULTS_DEFAULT, SEARCH_REGISTERED_MODEL_MAX_RESULTS_THRESHOLD, ) from mlflow.store.model_registry.abstract_store import AbstractStore from mlflow.store.model_registry.dbmodels.models import ( SqlModelVersion, SqlModelVersionTag, SqlRegisteredModel, SqlRegisteredModelAlias, SqlRegisteredModelTag, SqlWebhook, SqlWebhookEvent, ) from mlflow.tracking.client import MlflowClient from mlflow.utils.search_utils import SearchModelUtils, SearchModelVersionUtils, SearchUtils from mlflow.utils.time import get_current_time_millis from mlflow.utils.uri import extract_db_type_from_uri from mlflow.utils.validation import ( _REGISTERED_MODEL_ALIAS_LATEST, _validate_model_alias_name, _validate_model_alias_name_reserved, _validate_model_name, _validate_model_renaming, _validate_model_version, _validate_model_version_tag, _validate_registered_model_tag, _validate_tag_name, _validate_webhook_events, _validate_webhook_name, _validate_webhook_url, ) from mlflow.utils.workspace_utils import DEFAULT_WORKSPACE_NAME _logger = logging.getLogger(__name__) # Models that carry a ``workspace`` column and must be filtered # by the active workspace in every query. _WORKSPACE_MODELS = ( SqlRegisteredModel, SqlModelVersion, SqlWebhook, SqlRegisteredModelTag, SqlModelVersionTag, SqlRegisteredModelAlias, ) # For each database table, fetch its columns and define an appropriate attribute for each column # on the table's associated object representation (Mapper). This is necessary to ensure that # columns defined via backreference are available as Mapper instance attributes (e.g., # ``SqlRegisteredModel.model_versions``). For more information, see # https://docs.sqlalchemy.org/en/latest/orm/mapping_api.html#sqlalchemy.orm.configure_mappers # and https://docs.sqlalchemy.org/en/latest/orm/mapping_api.html#sqlalchemy.orm.mapper.Mapper sqlalchemy.orm.configure_mappers() class SqlAlchemyStore(AbstractStore): """ This entity may change or be removed in a future release without warning. SQLAlchemy compliant backend store for tracking meta data for MLflow entities. MLflow supports the database dialects ``mysql``, ``mssql``, ``sqlite``, and ``postgresql``. As specified in the `SQLAlchemy docs `_ , the database URI is expected in the format ``+://:@:/``. If you do not specify a driver, SQLAlchemy uses a dialect's default driver. This store interacts with SQL store using SQLAlchemy abstractions defined for MLflow entities. :py:class:`mlflow.store.model_registry.models.RegisteredModel` and :py:class:`mlflow.store.model_registry.models.ModelVersion` """ CREATE_MODEL_VERSION_RETRIES = 3 # Class-level cache for SQLAlchemy engines to prevent connection pool leaks # when multiple store instances are created with the same database URI. _engine_map: dict[str, sqlalchemy.engine.Engine] = {} _engine_map_lock = threading.Lock() @classmethod def _get_or_create_engine(cls, db_uri: str) -> sqlalchemy.engine.Engine: """Get a cached engine or create a new one for the given database URI.""" if db_uri not in cls._engine_map: with cls._engine_map_lock: if db_uri not in cls._engine_map: cls._engine_map[db_uri] = create_sqlalchemy_engine_with_retry(db_uri) return cls._engine_map[db_uri] def __init__(self, db_uri, read_db_uri=None): """ Create a database backed store. Args: db_uri: The SQLAlchemy database URI string to connect to the database. See the `SQLAlchemy docs `_ for format specifications. MLflow supports the dialects ``mysql``, ``mssql``, ``sqlite``, and ``postgresql``. read_db_uri: Optional SQLAlchemy database URI for a read replica. When provided, read operations are routed to this URI while write operations use ``db_uri``. If not provided, all operations use ``db_uri``. """ super().__init__() self.db_uri = db_uri self.db_type = extract_db_type_from_uri(db_uri) self.engine = self._get_or_create_engine(db_uri) if not _all_tables_exist(self.engine): _initialize_tables(self.engine) # Verify that all model registry tables exist. SqlAlchemyStore._verify_registry_tables_exist(self.engine) # Set up read replica engine if provided if read_db_uri and read_db_uri != db_uri: self.read_engine = self._get_or_create_engine(read_db_uri) WriteSessionMaker = sqlalchemy.orm.sessionmaker(bind=self.engine) ReadSessionMaker = sqlalchemy.orm.sessionmaker(bind=self.read_engine) self.ManagedSessionMaker = _get_routing_session_maker( WriteSessionMaker, ReadSessionMaker, self.db_type ) else: if read_db_uri and read_db_uri == db_uri: _logger.warning( "read_db_uri is the same as the primary db_uri; " "read replica routing will not be enabled. " "This is likely a configuration mistake." ) self.read_engine = None SessionMaker = sqlalchemy.orm.sessionmaker(bind=self.engine) self.ManagedSessionMaker = _get_managed_session_maker(SessionMaker, self.db_type) # TODO: verify schema here once we add logic to initialize the registry tables if they # don't exist (schema verification will fail in tests otherwise) # mlflow.store.db.utils._verify_schema(self.engine) self._initialize_store_state() @property def supports_workspaces(self) -> bool: """Indicates whether this store supports workspace isolation.""" return False def _get_active_workspace(self) -> str: """ Get the active workspace name. In single-tenant mode, always returns DEFAULT_WORKSPACE_NAME. Workspace-aware subclasses override this to enforce isolation. """ return DEFAULT_WORKSPACE_NAME def _get_query(self, session, model): """ Return a query for ``model``. Always filter on workspace for relevant models, to benefit from DB index. """ query = session.query(model) if model in _WORKSPACE_MODELS: query = query.filter(model.workspace == self._get_active_workspace()) return query def _with_workspace_field(self, instance): """ Allow subclasses to populate model fields (e.g., workspace metadata) on ORM instances. """ if hasattr(instance, "workspace") and getattr(instance, "workspace", None) is None: instance.workspace = DEFAULT_WORKSPACE_NAME return instance def _initialize_store_state(self): """ Initialize store state after construction. In single-tenant mode, validates no registry objects exist outside the default workspace. """ with self.ManagedSessionMaker() as session: exists_non_default_rm = ( session .query(SqlRegisteredModel.name) .filter(SqlRegisteredModel.workspace.isnot(None)) .filter(SqlRegisteredModel.workspace != DEFAULT_WORKSPACE_NAME) .first() is not None ) if exists_non_default_rm: name, workspace = ( session .query(SqlRegisteredModel.name, SqlRegisteredModel.workspace) .filter(SqlRegisteredModel.workspace != DEFAULT_WORKSPACE_NAME) .first() ) raise MlflowException( "Cannot disable workspaces because registered models exist outside the default " f"workspace (for example, model '{name}' in workspace '{workspace}'). " "Either remove those models or re-enable workspaces before starting.", error_code=INVALID_STATE, ) exists_non_default_webhook = ( session .query(SqlWebhook.webhook_id) .filter(SqlWebhook.workspace.isnot(None)) .filter(SqlWebhook.workspace != DEFAULT_WORKSPACE_NAME) .first() is not None ) if exists_non_default_webhook: webhook_id, workspace = ( session .query(SqlWebhook.webhook_id, SqlWebhook.workspace) .filter(SqlWebhook.workspace != DEFAULT_WORKSPACE_NAME) .first() ) raise MlflowException( "Cannot disable workspaces because webhooks exist outside the default " f"workspace (for example, webhook '{webhook_id}' in workspace '{workspace}'). " "Either remove those webhooks or re-enable workspaces before starting.", error_code=INVALID_STATE, ) def _get_dialect(self): return self.engine.dialect.name def _dispose_engine(self): self.engine.dispose() @staticmethod def _verify_registry_tables_exist(engine): # Verify that all tables have been created. inspected_tables = set(sqlalchemy.inspect(engine).get_table_names()) expected_tables = [ SqlRegisteredModel.__tablename__, SqlModelVersion.__tablename__, SqlWebhook.__tablename__, SqlWebhookEvent.__tablename__, ] if any(table not in inspected_tables for table in expected_tables): # TODO: Replace the MlflowException with the following line once it's possible to run # the registry against a different DB than the tracking server: # mlflow.store.db.utils._initialize_tables(self.engine) raise MlflowException("Database migration in unexpected state. Run manual upgrade.") @staticmethod def _get_eager_registered_model_query_options(): """ A list of SQLAlchemy query options that can be used to eagerly load the following registered model attributes when fetching a registered model: ``registered_model_tags`` and ``registered_model_aliases``. """ # Use a subquery load rather than a joined load in order to minimize the memory overhead # of the eager loading procedure. For more information about relationship loading # techniques, see https://docs.sqlalchemy.org/en/13/orm/ # loading_relationships.html#relationship-loading-techniques return [ sqlalchemy.orm.subqueryload(SqlRegisteredModel.registered_model_tags), sqlalchemy.orm.subqueryload(SqlRegisteredModel.registered_model_aliases), ] def _get_latest_versions_for_models( self, session, model_names: list[str] ) -> dict[str, list[SqlModelVersion]]: """ Batch-fetch the latest model version per stage for multiple registered models. Uses a SQL window function to compute the latest version per (name, stage) directly in the database, avoiding N+1 queries and Python-side iteration through all versions. """ if not model_names: return {} workspace_clauses = self._get_workspace_clauses(SqlModelVersion) row_num = ( sqlalchemy.func .row_number() .over( partition_by=[ SqlModelVersion.workspace, SqlModelVersion.name, SqlModelVersion.current_stage, ], order_by=SqlModelVersion.version.desc(), ) .label("rn") ) subquery = ( select(SqlModelVersion, row_num) .where( *workspace_clauses, SqlModelVersion.name.in_(model_names), SqlModelVersion.current_stage != STAGE_DELETED_INTERNAL, ) .subquery() ) query = ( select(SqlModelVersion) .where(*workspace_clauses) .join( subquery, sqlalchemy.and_( SqlModelVersion.workspace == subquery.c.workspace, SqlModelVersion.name == subquery.c.name, SqlModelVersion.version == subquery.c.version, ), ) .where(subquery.c.rn == 1) .options(*self._get_eager_model_version_query_options()) ) latest_versions = session.execute(query).scalars().all() result: dict[str, list[SqlModelVersion]] = {name: [] for name in model_names} for mv in latest_versions: result[mv.name].append(mv) return result @staticmethod def _get_eager_model_version_query_options(): """ A list of SQLAlchemy query options that can be used to eagerly load the following model version attributes when fetching a model version: ``model_version_tags``. """ # Use a subquery load rather than a joined load in order to minimize the memory overhead # of the eager loading procedure. For more information about relationship loading # techniques, see https://docs.sqlalchemy.org/en/13/orm/ # loading_relationships.html#relationship-loading-techniques return [sqlalchemy.orm.subqueryload(SqlModelVersion.model_version_tags)] def _get_workspace_clauses(self, model): """ Return workspace filter clauses for the model. Always filter on workspace for relevant models, to benefit from DB index. """ if model in _WORKSPACE_MODELS: return [model.workspace == self._get_active_workspace()] return [] def create_registered_model(self, name, tags=None, description=None, deployment_job_id=None): """ Create a new registered model in backend store. Args: name: Name of the new model. This is expected to be unique in the backend store. tags: A list of :py:class:`mlflow.entities.model_registry.RegisteredModelTag` instances associated with this registered model. description: Description of the version. deployment_job_id: Optional deployment job ID. Returns: A single object of :py:class:`mlflow.entities.model_registry.RegisteredModel` created in the backend. """ _validate_model_name(name) for tag in tags or []: _validate_registered_model_tag(tag.key, tag.value) with self.ManagedSessionMaker(read_only=False) as session: try: creation_time = get_current_time_millis() registered_model = self._with_workspace_field( SqlRegisteredModel( name=name, creation_time=creation_time, last_updated_time=creation_time, description=description, ) ) tags_dict = {} for tag in tags or []: tags_dict[tag.key] = tag.value registered_model.registered_model_tags = [ self._with_workspace_field( SqlRegisteredModelTag(name=name, key=key, value=value) ) for key, value in tags_dict.items() ] session.add(registered_model) session.flush() return registered_model.to_mlflow_entity() except sqlalchemy.exc.IntegrityError: existing_model = self.get_registered_model(name) handle_resource_already_exist_error( name, has_prompt_tag(existing_model._tags), has_prompt_tag(tags) ) def _get_registered_model(self, session, name, eager=False): """ Args: eager: If ``True``, eagerly loads the registered model's tags. If ``False``, these attributes are not eagerly loaded and will be loaded when their corresponding object properties are accessed from the resulting ``SqlRegisteredModel`` object. """ _validate_model_name(name) query = self._get_query(session, SqlRegisteredModel) if eager: query = query.options(*self._get_eager_registered_model_query_options()) rms = query.filter(SqlRegisteredModel.name == name).all() if len(rms) == 0: raise MlflowException( f"Registered Model with name={name} not found", RESOURCE_DOES_NOT_EXIST ) if len(rms) > 1: raise MlflowException( f"Expected only 1 registered model with name={name}. Found {len(rms)}.", INVALID_STATE, ) return rms[0] def update_registered_model(self, name, description, deployment_job_id=None): """ Update description of the registered model. Args: name: Registered model name. description: New description. deployment_job_id: Optional deployment job ID. Returns: A single updated :py:class:`mlflow.entities.model_registry.RegisteredModel` object. """ with self.ManagedSessionMaker(read_only=False) as session: sql_registered_model = self._get_registered_model(session, name) updated_time = get_current_time_millis() sql_registered_model.description = description sql_registered_model.last_updated_time = updated_time session.add(sql_registered_model) session.flush() return sql_registered_model.to_mlflow_entity() def rename_registered_model(self, name, new_name): """ Rename the registered model. Args: name: Registered model name. new_name: New proposed name. Returns: A single updated :py:class:`mlflow.entities.model_registry.RegisteredModel` object. """ _validate_model_renaming(new_name) with self.ManagedSessionMaker(read_only=False) as session: sql_registered_model = self._get_registered_model(session, name) try: updated_time = get_current_time_millis() sql_registered_model.name = new_name for sql_model_version in sql_registered_model.model_versions: sql_model_version.name = new_name sql_model_version.last_updated_time = updated_time sql_registered_model.last_updated_time = updated_time session.add_all([sql_registered_model] + sql_registered_model.model_versions) session.flush() return sql_registered_model.to_mlflow_entity() except sqlalchemy.exc.IntegrityError as e: raise MlflowException( f"Registered Model (name={new_name}) already exists. Error: {e}", RESOURCE_ALREADY_EXISTS, ) def delete_registered_model(self, name): """ Delete the registered model. Backend raises exception if a registered model with given name does not exist. Args: name: Registered model name. Returns: None """ with self.ManagedSessionMaker(read_only=False) as session: sql_registered_model = self._get_registered_model(session, name) session.delete(sql_registered_model) def _compute_next_token(self, max_results_for_query, current_size, offset, max_results): next_token = None if max_results_for_query == current_size: final_offset = offset + max_results next_token = SearchUtils.create_page_token(final_offset) return next_token def search_registered_models( self, filter_string=None, max_results=SEARCH_REGISTERED_MODEL_MAX_RESULTS_DEFAULT, order_by=None, page_token=None, ): """ Search for registered models in backend that satisfy the filter criteria. Args: filter_string: Filter query string, defaults to searching all registered models. max_results: Maximum number of registered models desired. order_by: List of column names with ASC|DESC annotation, to be used for ordering matching search results. page_token: Token specifying the next page of results. It should be obtained from a ``search_registered_models`` call. Returns: A PagedList of :py:class:`mlflow.entities.model_registry.RegisteredModel` objects that satisfy the search expressions. The pagination token for the next page can be obtained via the ``token`` attribute of the object. """ if max_results > SEARCH_REGISTERED_MODEL_MAX_RESULTS_THRESHOLD: raise MlflowException( "Invalid value for request parameter max_results. It must be at most " f"{SEARCH_REGISTERED_MODEL_MAX_RESULTS_THRESHOLD}, but got value {max_results}", INVALID_PARAMETER_VALUE, ) parsed_filters = SearchModelUtils.parse_search_filter(filter_string) parsed_orderby = self._parse_search_registered_models_order_by(order_by) offset = SearchUtils.parse_start_offset_from_page_token(page_token) # we query for max_results + 1 items to check whether there is another page to return. # this remediates having to make another query which returns no items. max_results_for_query = max_results + 1 with self.ManagedSessionMaker() as session: filter_query = self._get_search_registered_model_filter_query( session, parsed_filters, self.engine.dialect.name ) query = ( filter_query .options(*self._get_eager_registered_model_query_options()) .order_by(*parsed_orderby) .limit(max_results_for_query) ) if page_token: query = query.offset(offset) sql_registered_models = session.execute(query).scalars().all() next_page_token = self._compute_next_token( max_results_for_query, len(sql_registered_models), offset, max_results ) # Batch-fetch latest versions for all models to avoid N+1 queries model_names = [rm.name for rm in sql_registered_models[:max_results]] latest_versions_map = self._get_latest_versions_for_models(session, model_names) rm_entities = [ rm.to_mlflow_entity(preloaded_latest_versions=latest_versions_map.get(rm.name)) for rm in sql_registered_models[:max_results] ] return PagedList(rm_entities, next_page_token) def _get_search_registered_model_filter_query(self, session, parsed_filters, dialect): attribute_filters = [] tag_filters = {} tag_where_clauses = self._get_workspace_clauses(SqlRegisteredModelTag) for f in parsed_filters: type_ = f["type"] key = f["key"] comparator = f["comparator"] value = f["value"] if type_ == "attribute": if key != "name": raise MlflowException( f"Invalid attribute name: {key}", error_code=INVALID_PARAMETER_VALUE ) if comparator not in ("=", "!=", "LIKE", "ILIKE"): raise MlflowException( f"Invalid comparator for attribute: {comparator}", error_code=INVALID_PARAMETER_VALUE, ) attr = getattr(SqlRegisteredModel, key) attr_filter = SearchUtils.get_sql_comparison_func(comparator, dialect)(attr, value) attribute_filters.append(attr_filter) elif type_ == "tag": if comparator not in ("=", "!=", "LIKE", "ILIKE"): raise MlflowException.invalid_parameter_value( f"Invalid comparator for tag: {comparator}" ) if key not in tag_filters: key_filter = SearchUtils.get_sql_comparison_func("=", dialect)( SqlRegisteredModelTag.key, key ) tag_filters[key] = [key_filter] tag_filters[key].extend(tag_where_clauses) val_filter = SearchUtils.get_sql_comparison_func(comparator, dialect)( SqlRegisteredModelTag.value, value ) tag_filters[key].append(val_filter) else: raise MlflowException( f"Invalid token type: {type_}", error_code=INVALID_PARAMETER_VALUE ) attribute_filters.extend(self._get_workspace_clauses(SqlRegisteredModel)) rm_query = select(SqlRegisteredModel).filter(*attribute_filters) if not self._is_querying_prompt(parsed_filters): rm_query = self._update_query_to_exclude_prompts( rm_query, tag_filters, dialect, SqlRegisteredModel, SqlRegisteredModelTag, ) if tag_filters: sql_tag_filters = (sqlalchemy.and_(*x) for x in tag_filters.values()) tag_filter_query = ( select(SqlRegisteredModelTag.workspace, SqlRegisteredModelTag.name) .filter(sqlalchemy.or_(*sql_tag_filters)) .group_by(SqlRegisteredModelTag.workspace, SqlRegisteredModelTag.name) .having(sqlalchemy.func.count(sqlalchemy.literal(1)) == len(tag_filters)) .subquery() ) return rm_query.join( tag_filter_query, sqlalchemy.and_( SqlRegisteredModel.workspace == tag_filter_query.c.workspace, SqlRegisteredModel.name == tag_filter_query.c.name, ), ) else: return rm_query def _get_search_model_versions_filter_clauses(self, parsed_filters, dialect): attribute_filters = [] tag_filters = {} tag_where_clauses = self._get_workspace_clauses(SqlModelVersionTag) for f in parsed_filters: type_ = f["type"] key = f["key"] comparator = f["comparator"] value = f["value"] if type_ == "attribute": if key not in SearchModelVersionUtils.VALID_SEARCH_ATTRIBUTE_KEYS: raise MlflowException( f"Invalid attribute name: {key}", error_code=INVALID_PARAMETER_VALUE ) if key in SearchModelVersionUtils.NUMERIC_ATTRIBUTES: if ( comparator not in SearchModelVersionUtils.VALID_NUMERIC_ATTRIBUTE_COMPARATORS ): raise MlflowException( f"Invalid comparator for attribute {key}: {comparator}", error_code=INVALID_PARAMETER_VALUE, ) elif ( comparator not in SearchModelVersionUtils.VALID_STRING_ATTRIBUTE_COMPARATORS or (comparator == "IN" and key != "run_id") ): raise MlflowException( f"Invalid comparator for attribute: {comparator}", error_code=INVALID_PARAMETER_VALUE, ) if key == "source_path": key_name = "source" elif key == "version_number": key_name = "version" else: key_name = key attr = getattr(SqlModelVersion, key_name) if comparator == "IN": # Note: Here the run_id values in databases contain only lower case letters, # so we already filter out comparison values containing upper case letters # in `SearchModelUtils._get_value`. This addresses MySQL IN clause case # in-sensitive issue. val_filter = attr.in_(value) else: val_filter = SearchUtils.get_sql_comparison_func(comparator, dialect)( attr, value ) attribute_filters.append(val_filter) elif type_ == "tag": if comparator not in ("=", "!=", "LIKE", "ILIKE"): raise MlflowException.invalid_parameter_value( f"Invalid comparator for tag: {comparator}", ) if key not in tag_filters: key_filter = SearchUtils.get_sql_comparison_func("=", dialect)( SqlModelVersionTag.key, key ) tag_filters[key] = [key_filter] tag_filters[key].extend(tag_where_clauses) val_filter = SearchUtils.get_sql_comparison_func(comparator, dialect)( SqlModelVersionTag.value, value ) tag_filters[key].append(val_filter) else: raise MlflowException( f"Invalid token type: {type_}", error_code=INVALID_PARAMETER_VALUE ) attribute_filters.extend(self._get_workspace_clauses(SqlModelVersion)) mv_query = select(SqlModelVersion).filter(*attribute_filters) if not self._is_querying_prompt(parsed_filters): mv_query = self._update_query_to_exclude_prompts( mv_query, tag_filters, dialect, SqlModelVersion, SqlModelVersionTag, ) if tag_filters: sql_tag_filters = (sqlalchemy.and_(*x) for x in tag_filters.values()) tag_filter_query = ( select( SqlModelVersionTag.workspace, SqlModelVersionTag.name, SqlModelVersionTag.version, ) .filter(sqlalchemy.or_(*sql_tag_filters)) .group_by( SqlModelVersionTag.workspace, SqlModelVersionTag.name, SqlModelVersionTag.version, ) .having(sqlalchemy.func.count(sqlalchemy.literal(1)) == len(tag_filters)) .subquery() ) return mv_query.join( tag_filter_query, sqlalchemy.and_( SqlModelVersion.workspace == tag_filter_query.c.workspace, SqlModelVersion.name == tag_filter_query.c.name, SqlModelVersion.version == tag_filter_query.c.version, ), ) else: return mv_query def _update_query_to_exclude_prompts( self, query: Any, tag_filters: dict[str, list[Any]], dialect: str, main_db_model: SqlModelVersion | SqlRegisteredModel, tag_db_model: SqlModelVersionTag | SqlRegisteredModelTag, ): """ Update query to exclude all prompt rows and return only normal model or model versions. Prompts and normal models are distinguished by the `mlflow.prompt.is_prompt` tag. The search API should only return normal models by default. However, simply filtering rows using the tag like this does not work because models do not have the prompt tag. tags.`mlflow.prompt.is_prompt` != 'true' tags.`mlflow.prompt.is_prompt` = 'false' To workaround this, we need to use a subquery to get all prompt rows and then use an anti-join for excluding prompts. """ # If the tag filter contains the prompt tag, remove it tag_filters.pop(IS_PROMPT_TAG_KEY, []) # Filter to get all prompt rows equal = SearchUtils.get_sql_comparison_func("=", dialect) prompts_subquery = ( select(tag_db_model.workspace, tag_db_model.name) .filter( equal(tag_db_model.key, IS_PROMPT_TAG_KEY), equal(tag_db_model.value, "true"), *self._get_workspace_clauses(tag_db_model), ) .group_by(tag_db_model.workspace, tag_db_model.name) .subquery() ) return query.join( prompts_subquery, sqlalchemy.and_( main_db_model.workspace == prompts_subquery.c.workspace, main_db_model.name == prompts_subquery.c.name, ), isouter=True, ).filter(prompts_subquery.c.name.is_(None)) @classmethod def _is_querying_prompt(cls, parsed_filters: list[dict[str, Any]]) -> bool: for f in parsed_filters: if f["type"] != "tag" or f["key"] != IS_PROMPT_TAG_KEY: continue return (f["comparator"] == "=" and f["value"].lower() == "true") or ( f["comparator"] == "!=" and f["value"].lower() == "false" ) # Query should return only normal models by default return False @classmethod def _parse_search_registered_models_order_by(cls, order_by_list): """Sorts a set of registered models based on their natural ordering and an overriding set of order_bys. Registered models are naturally ordered first by name ascending. """ clauses = [] observed_order_by_clauses = set() if order_by_list: for order_by_clause in order_by_list: ( attribute_token, ascending, ) = SearchUtils.parse_order_by_for_search_registered_models(order_by_clause) if attribute_token == SqlRegisteredModel.name.key: field = SqlRegisteredModel.name elif attribute_token in SearchUtils.VALID_TIMESTAMP_ORDER_BY_KEYS: field = SqlRegisteredModel.last_updated_time else: raise MlflowException( f"Invalid order by key '{attribute_token}' specified." + "Valid keys are " + f"'{SearchUtils.RECOMMENDED_ORDER_BY_KEYS_REGISTERED_MODELS}'", error_code=INVALID_PARAMETER_VALUE, ) if field.key in observed_order_by_clauses: raise MlflowException(f"`order_by` contains duplicate fields: {order_by_list}") observed_order_by_clauses.add(field.key) if ascending: clauses.append(field.asc()) else: clauses.append(field.desc()) if SqlRegisteredModel.name.key not in observed_order_by_clauses: clauses.append(SqlRegisteredModel.name.asc()) return clauses def get_registered_model(self, name): """ Get registered model instance by name. Args: name: Registered model name. Returns: A single :py:class:`mlflow.entities.model_registry.RegisteredModel` object. """ with self.ManagedSessionMaker() as session: return self._get_registered_model(session, name, eager=True).to_mlflow_entity() def get_latest_versions(self, name, stages=None): """ Latest version models for each requested stage. If no ``stages`` argument is provided, returns the latest version for each stage. Args: name: Registered model name. stages: List of desired stages. If input list is None, return latest versions for each stage. Returns: List of :py:class:`mlflow.entities.model_registry.ModelVersion` objects. """ with self.ManagedSessionMaker() as session: sql_registered_model = self._get_registered_model(session, name) # Convert to RegisteredModel entity first and then extract latest_versions latest_versions = sql_registered_model.to_mlflow_entity().latest_versions if stages is None or len(stages) == 0: expected_stages = {get_canonical_stage(stage) for stage in ALL_STAGES} else: expected_stages = {get_canonical_stage(stage) for stage in stages} mvs = [mv for mv in latest_versions if mv.current_stage in expected_stages] # Populate aliases for each model version for mv in mvs: model_aliases = sql_registered_model.registered_model_aliases mv.aliases = [alias.alias for alias in model_aliases if alias.version == mv.version] return mvs def _get_registered_model_tag(self, session, name, key): tags = ( self ._get_query(session, SqlRegisteredModelTag) .filter( SqlRegisteredModelTag.name == name, SqlRegisteredModelTag.key == key, ) .all() ) if len(tags) == 0: return None if len(tags) > 1: raise MlflowException( f"Expected only 1 registered model tag with name={name}, key={key}. " f"Found {len(tags)}.", INVALID_STATE, ) return tags[0] def set_registered_model_tag(self, name, tag): """ Set a tag for the registered model. Args: name: Registered model name. tag: :py:class:`mlflow.entities.model_registry.RegisteredModelTag` instance to log. Returns: None """ _validate_model_name(name) _validate_registered_model_tag(tag.key, tag.value) with self.ManagedSessionMaker(read_only=False) as session: # check if registered model exists sql_registered_model = self._get_registered_model(session, name) session.merge( SqlRegisteredModelTag( workspace=sql_registered_model.workspace, name=name, key=tag.key, value=tag.value, ) ) def delete_registered_model_tag(self, name, key): """ Delete a tag associated with the registered model. Args: name: Registered model name. key: Registered model tag key. Returns: None """ _validate_model_name(name) _validate_tag_name(key) with self.ManagedSessionMaker(read_only=False) as session: # check if registered model exists self._get_registered_model(session, name) existing_tag = self._get_registered_model_tag(session, name, key) if existing_tag is not None: session.delete(existing_tag) # CRUD API for ModelVersion objects def create_model_version( self, name, source, run_id=None, tags=None, run_link=None, description=None, local_model_path=None, model_id: str | None = None, ): """ Create a new model version from given source and run ID. Args: name: Registered model name. source: URI indicating the location of the model artifacts. run_id: Run ID from MLflow tracking server that generated the model. tags: A list of :py:class:`mlflow.entities.model_registry.ModelVersionTag` instances associated with this model version. run_link: Link to the run from an MLflow tracking server that generated this model. description: Description of the version. local_model_path: Unused. model_id: The ID of the model (from an Experiment) that is being promoted to a registered model version, if applicable. Returns: A single object of :py:class:`mlflow.entities.model_registry.ModelVersion` created in the backend. """ _validate_model_name(name) for tag in tags or []: _validate_model_version_tag(tag.key, tag.value) storage_location = source if urllib.parse.urlparse(source).scheme == "models": parsed_model_uri = _parse_model_uri(source) try: if parsed_model_uri.model_id is not None: # TODO: Propagate tracking URI to file sqlalchemy directly, rather than relying # on global URI (individual MlflowClient instances may have different tracking # URIs) model = MlflowClient().get_logged_model(parsed_model_uri.model_id) storage_location = model.artifact_location run_id = run_id or model.source_run_id else: storage_location = self.get_model_version_download_uri( parsed_model_uri.name, parsed_model_uri.version ) except Exception as e: raise MlflowException( f"Unable to fetch model from model URI source artifact location '{source}'." f"Error: {e}" ) from e if not run_id and model_id: model = MlflowClient().get_logged_model(model_id) run_id = model.source_run_id with self.ManagedSessionMaker(read_only=False) as session: creation_time = get_current_time_millis() for attempt in range(self.CREATE_MODEL_VERSION_RETRIES): try: sql_registered_model = self._get_registered_model(session, name) sql_registered_model.last_updated_time = creation_time max_version = ( session .query(sqlalchemy.func.max(SqlModelVersion.version)) .filter( SqlModelVersion.name == name, *self._get_workspace_clauses(SqlModelVersion), ) .scalar() ) version = (max_version or 0) + 1 model_version = self._with_workspace_field( SqlModelVersion( name=name, version=version, creation_time=creation_time, last_updated_time=creation_time, source=source, storage_location=storage_location, run_id=run_id, run_link=run_link, description=description, ) ) tags_dict = {} for tag in tags or []: tags_dict[tag.key] = tag.value model_version.model_version_tags = [ self._with_workspace_field( SqlModelVersionTag(name=name, version=version, key=key, value=value) ) for key, value in tags_dict.items() ] session.add_all([sql_registered_model, model_version]) session.flush() return self._populate_model_version_aliases( session, name, model_version.to_mlflow_entity() ) except sqlalchemy.exc.IntegrityError: session.rollback() more_retries = self.CREATE_MODEL_VERSION_RETRIES - attempt - 1 _logger.info( "Model Version creation error (name=%s) Retrying %s more time%s.", name, str(more_retries), "s" if more_retries > 1 else "", ) raise MlflowException( f"Model Version creation error (name={name}). Giving up after " f"{self.CREATE_MODEL_VERSION_RETRIES} attempts." ) def _populate_model_version_aliases(self, session, name, version): model_aliases = self._get_registered_model(session, name).registered_model_aliases version.aliases = [ alias.alias for alias in model_aliases if alias.version == version.version ] return version def _get_model_version_from_db(self, session, name, version, conditions, query_options=None): if query_options is None: query_options = [] versions = ( self ._get_query(session, SqlModelVersion) .options(*query_options) .filter(*conditions) .all() ) if len(versions) == 0: raise MlflowException( f"Model Version (name={name}, version={version}) not found", RESOURCE_DOES_NOT_EXIST, ) if len(versions) > 1: raise MlflowException( f"Expected only 1 model version with (name={name}, version={version}). " f"Found {len(versions)}.", INVALID_STATE, ) return versions[0] def _get_sql_model_version(self, session, name, version, eager=False): """ Args: eager: If ``True``, eagerly loads the model version's tags. If ``False``, these attributes are not eagerly loaded and will be loaded when their corresponding object properties are accessed from the resulting ``SqlModelVersion`` object. """ _validate_model_name(name) _validate_model_version(version) query_options = self._get_eager_model_version_query_options() if eager else [] conditions = [ SqlModelVersion.name == name, SqlModelVersion.version == version, SqlModelVersion.current_stage != STAGE_DELETED_INTERNAL, ] return self._get_model_version_from_db(session, name, version, conditions, query_options) def _get_sql_model_version_including_deleted(self, name, version): """ Private method to retrieve model versions including those that are internally deleted. Used in tests to verify redaction behavior on deletion. Args: name: Registered model name. version: Registered model version. Returns: A single :py:class:`mlflow.entities.model_registry.ModelVersion` object. """ with self.ManagedSessionMaker() as session: conditions = [ SqlModelVersion.name == name, SqlModelVersion.version == version, ] sql_model_version = self._get_model_version_from_db(session, name, version, conditions) return self._populate_model_version_aliases( session, name, sql_model_version.to_mlflow_entity() ) def update_model_version(self, name, version, description=None): """ Update metadata associated with a model version in backend. Args: name: Registered model name. version: Registered model version. description: New model description. Returns: A single :py:class:`mlflow.entities.model_registry.ModelVersion` object. """ with self.ManagedSessionMaker(read_only=False) as session: updated_time = get_current_time_millis() sql_model_version = self._get_sql_model_version(session, name=name, version=version) sql_model_version.description = description sql_model_version.last_updated_time = updated_time session.add(sql_model_version) return self._populate_model_version_aliases( session, name, sql_model_version.to_mlflow_entity() ) def transition_model_version_stage(self, name, version, stage, archive_existing_versions): """ Update model version stage. Args: name: Registered model name. version: Registered model version. stage: New desired stage for this model version. archive_existing_versions: If this flag is set to ``True``, all existing model versions in the stage will be automatically moved to the "archived" stage. Only valid when ``stage`` is ``"staging"`` or ``"production"`` otherwise an error will be raised. Returns: A single :py:class:`mlflow.entities.model_registry.ModelVersion` object. """ is_active_stage = get_canonical_stage(stage) in DEFAULT_STAGES_FOR_GET_LATEST_VERSIONS if archive_existing_versions and not is_active_stage: msg_tpl = ( "Model version transition cannot archive existing model versions " "because '{}' is not an Active stage. Valid stages are {}" ) raise MlflowException(msg_tpl.format(stage, DEFAULT_STAGES_FOR_GET_LATEST_VERSIONS)) with self.ManagedSessionMaker(read_only=False) as session: last_updated_time = get_current_time_millis() model_versions = [] if archive_existing_versions: conditions = [ SqlModelVersion.name == name, SqlModelVersion.version != version, SqlModelVersion.current_stage == get_canonical_stage(stage), ] model_versions = self._get_query(session, SqlModelVersion).filter(*conditions).all() for mv in model_versions: mv.current_stage = STAGE_ARCHIVED mv.last_updated_time = last_updated_time sql_model_version = self._get_sql_model_version( session=session, name=name, version=version ) sql_model_version.current_stage = get_canonical_stage(stage) sql_model_version.last_updated_time = last_updated_time sql_registered_model = sql_model_version.registered_model sql_registered_model.last_updated_time = last_updated_time session.add_all([*model_versions, sql_model_version, sql_registered_model]) return self._populate_model_version_aliases( session, name, sql_model_version.to_mlflow_entity() ) def delete_model_version(self, name, version): """ Delete model version in backend. Args: name: Registered model name. version: Registered model version. Returns: None """ # currently delete model version still keeps the tags associated with the version with self.ManagedSessionMaker(read_only=False) as session: updated_time = get_current_time_millis() sql_model_version = self._get_sql_model_version(session, name, version) sql_registered_model = sql_model_version.registered_model sql_registered_model.last_updated_time = updated_time aliases = sql_registered_model.registered_model_aliases for alias in aliases: if alias.version == version: session.delete(alias) sql_model_version.current_stage = STAGE_DELETED_INTERNAL sql_model_version.last_updated_time = updated_time sql_model_version.description = None sql_model_version.user_id = None sql_model_version.source = "REDACTED-SOURCE-PATH" sql_model_version.run_id = "REDACTED-RUN-ID" sql_model_version.run_link = "REDACTED-RUN-LINK" sql_model_version.status_message = None session.add_all([sql_registered_model, sql_model_version]) def get_model_version(self, name, version): """ Get the model version instance by name and version. Args: name: Registered model name. version: Registered model version. Returns: A single :py:class:`mlflow.entities.model_registry.ModelVersion` object. """ with self.ManagedSessionMaker() as session: sql_model_version = self._get_sql_model_version(session, name, version, eager=True) return self._populate_model_version_aliases( session, name, sql_model_version.to_mlflow_entity() ) def get_model_version_download_uri(self, name, version): """ Get the download location in Model Registry for this model version. NOTE: For first version of Model Registry, since the models are not copied over to another location, download URI points to input source path. Args: name: Registered model name. version: Registered model version. Returns: A single URI location that allows reads for downloading. """ with self.ManagedSessionMaker() as session: sql_model_version = self._get_sql_model_version(session, name, version) return sql_model_version.storage_location or sql_model_version.source def search_model_versions( self, filter_string=None, max_results=SEARCH_MODEL_VERSION_MAX_RESULTS_DEFAULT, order_by=None, page_token=None, ): """ Search for model versions in backend that satisfy the filter criteria. Args: filter_string: A filter string expression. Currently supports a single filter condition either name of model like ``name = 'model_name'`` or ``run_id = '...'``. max_results: Maximum number of model versions desired. order_by: List of column names with ASC|DESC annotation, to be used for ordering matching search results. page_token: Token specifying the next page of results. It should be obtained from a ``search_model_versions`` call. Returns: A PagedList of :py:class:`mlflow.entities.model_registry.ModelVersion` objects that satisfy the search expressions. The pagination token for the next page can be obtained via the ``token`` attribute of the object. """ if not isinstance(max_results, int) or max_results < 1: raise MlflowException( "Invalid value for max_results. It must be a positive integer," f" but got {max_results}", INVALID_PARAMETER_VALUE, ) if max_results > SEARCH_MODEL_VERSION_MAX_RESULTS_THRESHOLD: raise MlflowException( "Invalid value for request parameter max_results. It must be at most " f"{SEARCH_MODEL_VERSION_MAX_RESULTS_THRESHOLD}, but got value {max_results}", INVALID_PARAMETER_VALUE, ) parsed_filters = SearchModelVersionUtils.parse_search_filter(filter_string) filter_query = self._get_search_model_versions_filter_clauses( parsed_filters, self.engine.dialect.name ) parsed_orderby = self._parse_search_model_versions_order_by( order_by or ["last_updated_timestamp DESC", "name ASC", "version_number DESC"] ) offset = SearchUtils.parse_start_offset_from_page_token(page_token) # we query for max_results + 1 items to check whether there is another page to return. # this remediates having to make another query which returns no items. max_results_for_query = max_results + 1 with self.ManagedSessionMaker() as session: query = ( filter_query .options(*self._get_eager_model_version_query_options()) .filter(SqlModelVersion.current_stage != STAGE_DELETED_INTERNAL) .order_by(*parsed_orderby) .limit(max_results_for_query) ) if page_token: query = query.offset(offset) sql_model_versions = session.execute(query).scalars().all() next_page_token = self._compute_next_token( max_results_for_query, len(sql_model_versions), offset, max_results ) model_versions = [mv.to_mlflow_entity() for mv in sql_model_versions][:max_results] return PagedList(model_versions, next_page_token) @classmethod def _parse_search_model_versions_order_by(cls, order_by_list): """Sorts a set of model versions based on their natural ordering and an overriding set of order_bys. Model versions are naturally ordered first by name ascending, then by version ascending. """ clauses = [] observed_order_by_clauses = set() if order_by_list: for order_by_clause in order_by_list: ( _, key, ascending, ) = SearchModelVersionUtils.parse_order_by_for_search_model_versions( order_by_clause ) if key not in SearchModelVersionUtils.VALID_ORDER_BY_ATTRIBUTE_KEYS: raise MlflowException( f"Invalid order by key '{key}' specified. " "Valid keys are " f"{SearchModelVersionUtils.VALID_ORDER_BY_ATTRIBUTE_KEYS}", error_code=INVALID_PARAMETER_VALUE, ) else: if key == "version_number": field = SqlModelVersion.version elif key == "creation_timestamp": field = SqlModelVersion.creation_time elif key == "last_updated_timestamp": field = SqlModelVersion.last_updated_time else: field = getattr(SqlModelVersion, key) if field.key in observed_order_by_clauses: raise MlflowException(f"`order_by` contains duplicate fields: {order_by_list}") observed_order_by_clauses.add(field.key) if ascending: clauses.append(field.asc()) else: clauses.append(field.desc()) if SqlModelVersion.name.key not in observed_order_by_clauses: clauses.append(SqlModelVersion.name.asc()) if SqlModelVersion.version.key not in observed_order_by_clauses: clauses.append(SqlModelVersion.version.desc()) return clauses def _get_model_version_tag(self, session, name, version, key): tags = ( self ._get_query(session, SqlModelVersionTag) .filter( SqlModelVersionTag.name == name, SqlModelVersionTag.version == version, SqlModelVersionTag.key == key, ) .all() ) if len(tags) == 0: return None if len(tags) > 1: raise MlflowException( f"Expected only 1 model version tag with name={name}, version={version}, " f"key={key}. Found {len(tags)}.", INVALID_STATE, ) return tags[0] def set_model_version_tag(self, name, version, tag): """ Set a tag for the model version. Args: name: Registered model name. version: Registered model version. tag: :py:class:`mlflow.entities.model_registry.ModelVersionTag` instance to log. Returns: None """ _validate_model_name(name) _validate_model_version(version) _validate_model_version_tag(tag.key, tag.value) with self.ManagedSessionMaker(read_only=False) as session: # check if model version exists sql_model_version = self._get_sql_model_version(session, name, version) session.merge( SqlModelVersionTag( workspace=sql_model_version.workspace, name=name, version=version, key=tag.key, value=tag.value, ) ) def delete_model_version_tag(self, name, version, key): """ Delete a tag associated with the model version. Args: name: Registered model name. version: Registered model version. key: Tag key. Returns: None """ _validate_model_name(name) _validate_model_version(version) _validate_tag_name(key) with self.ManagedSessionMaker(read_only=False) as session: # check if model version exists self._get_sql_model_version(session, name, version) existing_tag = self._get_model_version_tag(session, name, version, key) if existing_tag is not None: session.delete(existing_tag) def _get_registered_model_alias(self, session, name, alias): return ( self ._get_query(session, SqlRegisteredModelAlias) .filter( SqlRegisteredModelAlias.name == name, SqlRegisteredModelAlias.alias == alias, ) .first() ) def set_registered_model_alias(self, name, alias, version): """ Set a registered model alias pointing to a model version. Args: name: Registered model name. alias: Name of the alias. version: Registered model version number. Returns: None """ _validate_model_name(name) _validate_model_alias_name(alias) _validate_model_alias_name_reserved(alias) _validate_model_version(version) with self.ManagedSessionMaker(read_only=False) as session: # check if model version exists sql_model_version = self._get_sql_model_version(session, name, version) session.merge( SqlRegisteredModelAlias( workspace=sql_model_version.workspace, name=name, alias=alias, version=version, ) ) def delete_registered_model_alias(self, name, alias): """ Delete an alias associated with a registered model. Args: name: Registered model name. alias: Name of the alias. Returns: None """ _validate_model_name(name) _validate_model_alias_name(alias) with self.ManagedSessionMaker(read_only=False) as session: # check if registered model exists self._get_registered_model(session, name) existing_alias = self._get_registered_model_alias(session, name, alias) if existing_alias is not None: session.delete(existing_alias) def get_model_version_by_alias(self, name, alias): """ Get the model version instance by name and alias. Args: name: Registered model name. alias: Name of the alias. Returns: A single :py:class:`mlflow.entities.model_registry.ModelVersion` object. """ _validate_model_name(name) _validate_model_alias_name(alias) if alias.lower() == _REGISTERED_MODEL_ALIAS_LATEST: if versions := self.get_latest_versions(name): return versions[0] else: raise MlflowException( f"Latest version not found for model {name}.", RESOURCE_DOES_NOT_EXIST ) with self.ManagedSessionMaker() as session: # check if registered model exists self._get_registered_model(session, name) existing_alias = self._get_registered_model_alias(session, name, alias) if existing_alias is not None: sql_model_version = self._get_sql_model_version( session, existing_alias.name, existing_alias.version ) return self._populate_model_version_aliases( session, name, sql_model_version.to_mlflow_entity() ) else: raise MlflowException( f"Registered model alias {alias} not found.", INVALID_PARAMETER_VALUE ) def _await_model_version_creation(self, mv, await_creation_for): """ Does not wait for the model version to become READY as a successful creation will immediately place the model version in a READY state. """ # Webhook CRUD operations def create_webhook( self, name: str, url: str, events: list[WebhookEvent], description: str | None = None, secret: str | None = None, status: WebhookStatus | None = None, ) -> Webhook: _validate_webhook_name(name) _validate_webhook_url(url) _validate_webhook_events(events) with self.ManagedSessionMaker(read_only=False) as session: webhook_id = str(uuid.uuid4()) creation_time = get_current_time_millis() webhook = self._with_workspace_field( SqlWebhook( webhook_id=webhook_id, name=name, url=url, description=description, secret=secret, status=(status or WebhookStatus.ACTIVE).value, creation_timestamp=creation_time, last_updated_timestamp=creation_time, ) ) session.add(webhook) session.add_all( SqlWebhookEvent( webhook_id=webhook_id, entity=e.entity.value, action=e.action.value, ) for e in events ) session.flush() return webhook.to_mlflow_entity() def get_webhook(self, webhook_id: str) -> Webhook: with self.ManagedSessionMaker() as session: webhook = self._get_webhook_by_id(session, webhook_id) return webhook.to_mlflow_entity() def list_webhooks( self, max_results: int | None = None, page_token: str | None = None, ) -> PagedList[Webhook]: max_results = max_results or 100 if max_results < 1 or max_results > 1000: raise MlflowException( "max_results must be between 1 and 1000.", INVALID_PARAMETER_VALUE ) offset = SearchUtils.parse_start_offset_from_page_token(page_token) with self.ManagedSessionMaker() as session: query = ( self ._get_query(session, SqlWebhook) .filter(SqlWebhook.deleted_timestamp.is_(None)) .order_by(SqlWebhook.creation_timestamp.desc()) .limit(max_results + 1) ) if page_token: query = query.offset(offset) webhooks = query.all() # Check if there's a next page has_next_page = len(webhooks) > max_results next_page_token = None if has_next_page: webhooks = webhooks[:max_results] next_page_token = SearchUtils.create_page_token(offset + max_results) return PagedList([w.to_mlflow_entity() for w in webhooks], next_page_token) def list_webhooks_by_event( self, event: WebhookEvent, max_results: int | None = None, page_token: str | None = None, ) -> PagedList[Webhook]: max_results = max_results or 100 if max_results < 1 or max_results > 1000: raise MlflowException( "max_results must be between 1 and 1000.", INVALID_PARAMETER_VALUE ) offset = SearchUtils.parse_start_offset_from_page_token(page_token) with self.ManagedSessionMaker() as session: # Query webhooks that have the specific event in their related webhook_events query = ( self ._get_query(session, SqlWebhook) .join(SqlWebhookEvent) .filter(SqlWebhook.deleted_timestamp.is_(None)) .filter(SqlWebhookEvent.entity == event.entity.value) .filter(SqlWebhookEvent.action == event.action.value) .order_by(SqlWebhook.creation_timestamp.desc()) .limit(max_results + 1) ) if page_token: query = query.offset(offset) webhooks = query.all() # Check if there's a next page has_next_page = len(webhooks) > max_results next_page_token = None if has_next_page: webhooks = webhooks[:max_results] next_page_token = SearchUtils.create_page_token(offset + max_results) return PagedList([w.to_mlflow_entity() for w in webhooks], next_page_token) def update_webhook( self, webhook_id: str, name: str | None = None, description: str | None = None, url: str | None = None, events: list[WebhookEvent] | None = None, secret: str | None = None, status: WebhookStatus | None = None, ) -> Webhook: with self.ManagedSessionMaker(read_only=False) as session: webhook = self._get_webhook_by_id(session, webhook_id) # Update fields if provided if name is not None: _validate_webhook_name(name) webhook.name = name if url is not None: _validate_webhook_url(url) webhook.url = url if events is not None: _validate_webhook_events(events) # Delete existing webhook events session.query(SqlWebhookEvent).filter( SqlWebhookEvent.webhook_id == webhook_id ).delete() # Create new webhook events session.add_all( SqlWebhookEvent( webhook_id=webhook_id, entity=e.entity.value, action=e.action.value ) for e in events ) if description is not None: webhook.description = description if secret is not None: webhook.secret = secret if status is not None: webhook.status = status.value webhook.last_updated_timestamp = get_current_time_millis() session.add(webhook) session.flush() return webhook.to_mlflow_entity() def delete_webhook(self, webhook_id: str) -> None: with self.ManagedSessionMaker(read_only=False) as session: webhook = self._get_webhook_by_id(session, webhook_id) # Soft delete by setting deleted_timestamp webhook.deleted_timestamp = get_current_time_millis() webhook.last_updated_timestamp = webhook.deleted_timestamp session.add(webhook) session.flush() # Helper methods for webhooks def _get_webhook_by_id(self, session: Session, webhook_id: str) -> SqlWebhook: if webhook := ( self ._get_query(session, SqlWebhook) .filter( SqlWebhook.webhook_id == webhook_id, SqlWebhook.deleted_timestamp.is_(None), ) .first() ): return webhook raise MlflowException(f"Webhook with ID {webhook_id} not found.", RESOURCE_DOES_NOT_EXIST)