项目文件夹

文件
2026-07-13 13:22:34 +08:00

1794 行
72 KiB
Python

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 <https://docs.sqlalchemy.org/en/latest/core/engines.html#database-urls>`_ ,
the database URI is expected in the format
``<dialect>+<driver>://<username>:<password>@<host>:<port>/<database>``. 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
<https://docs.sqlalchemy.org/en/latest/core/engines.html#database-urls>`_
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)