This commit is contained in:
Christian Mantha
2026-03-02 19:10:52 -05:00
commit 2ca0b9ef7c
28907 changed files with 5233713 additions and 0 deletions

View File

@@ -0,0 +1,36 @@
from abc import ABCMeta, abstractmethod
from mlflow.utils.annotations import developer_stable
@developer_stable
class RequestHeaderProvider:
"""
Abstract base class for specifying custom request headers to add to outgoing requests
(e.g. request headers specifying the environment from which mlflow is running).
When a request is sent, MLflow will iterate through all registered RequestHeaderProviders.
For each provider where ``in_context`` returns ``True``, MLflow calls the ``request_headers``
method on the provider to compute request headers.
All resulting request headers will then be merged together and sent with the request.
"""
__metaclass__ = ABCMeta
@abstractmethod
def in_context(self):
"""Determine if MLflow is running in this context.
Returns:
bool indicating if in this context.
"""
@abstractmethod
def request_headers(self):
"""Generate context-specific request headers.
Returns:
dict of request headers.
"""

View File

@@ -0,0 +1,38 @@
from mlflow.tracking.request_header.abstract_request_header_provider import RequestHeaderProvider
from mlflow.utils import databricks_utils
class DatabricksRequestHeaderProvider(RequestHeaderProvider):
"""
Provides request headers indicating the type of Databricks environment from which a request
was made.
"""
def in_context(self):
return (
databricks_utils.is_in_cluster()
or databricks_utils.is_in_databricks_notebook()
or databricks_utils.is_in_databricks_job()
)
def request_headers(self):
request_headers = {}
if databricks_utils.is_in_databricks_notebook():
request_headers["notebook_id"] = databricks_utils.get_notebook_id()
if databricks_utils.is_in_databricks_job():
request_headers["job_id"] = databricks_utils.get_job_id()
request_headers["job_run_id"] = databricks_utils.get_job_run_id()
request_headers["job_type"] = databricks_utils.get_job_type()
if databricks_utils.is_in_cluster():
request_headers["cluster_id"] = databricks_utils.get_cluster_id()
command_run_id = databricks_utils.get_command_run_id()
if command_run_id is not None:
request_headers["command_run_id"] = command_run_id
workload_id = databricks_utils.get_workload_id()
workload_class = databricks_utils.get_workload_class()
if workload_id is not None:
request_headers["workload_id"] = workload_id
if workload_class is not None:
request_headers["workload_class"] = workload_class
return request_headers

View File

@@ -0,0 +1,17 @@
from mlflow import __version__
from mlflow.tracking.request_header.abstract_request_header_provider import RequestHeaderProvider
_USER_AGENT = "User-Agent"
_DEFAULT_HEADERS = {_USER_AGENT: f"mlflow-python-client/{__version__}"}
class DefaultRequestHeaderProvider(RequestHeaderProvider):
"""
Provides default request headers for outgoing request.
"""
def in_context(self):
return True
def request_headers(self):
return dict(**_DEFAULT_HEADERS)

View File

@@ -0,0 +1,79 @@
import logging
import warnings
from mlflow.tracking.request_header.databricks_request_header_provider import (
DatabricksRequestHeaderProvider,
)
from mlflow.tracking.request_header.default_request_header_provider import (
DefaultRequestHeaderProvider,
)
from mlflow.utils.plugins import get_entry_points
_logger = logging.getLogger(__name__)
class RequestHeaderProviderRegistry:
def __init__(self):
self._registry = []
def register(self, request_header_provider):
self._registry.append(request_header_provider())
def register_entrypoints(self):
"""Register tracking stores provided by other packages"""
for entrypoint in get_entry_points("mlflow.request_header_provider"):
try:
self.register(entrypoint.load())
except (AttributeError, ImportError) as exc:
warnings.warn(
'Failure attempting to register request header provider "{}": {}'.format(
entrypoint.name, str(exc)
),
stacklevel=2,
)
def __iter__(self):
return iter(self._registry)
_request_header_provider_registry = RequestHeaderProviderRegistry()
_request_header_provider_registry.register(DatabricksRequestHeaderProvider)
_request_header_provider_registry.register(DefaultRequestHeaderProvider)
_request_header_provider_registry.register_entrypoints()
def resolve_request_headers(request_headers=None):
"""Generate a set of request headers from registered providers.
Request headers are resolved in the order that providers are registered. Argument headers are
applied last. This function iterates through all request header providers in the registry.
Additional context providers can be registered as described in
:py:class:`mlflow.tracking.request_header.RequestHeaderProvider`.
Args:
request_headers: A dictionary of request headers to override. If specified, headers passed
in this argument will override those inferred from the context.
Returns:
A dictionary of resolved headers.
"""
all_request_headers = {}
for provider in _request_header_provider_registry:
try:
if provider.in_context():
# all_request_headers.update(provider.request_headers())
for header, value in provider.request_headers().items():
all_request_headers[header] = (
f"{all_request_headers[header]} {value}"
if header in all_request_headers
else value
)
except Exception as e:
_logger.warning("Encountered unexpected error during resolving request headers: %s", e)
if request_headers is not None:
all_request_headers.update(request_headers)
return all_request_headers