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,17 @@
# A special tag in RegisteredModel to indicate that it is a prompt
import re
IS_PROMPT_TAG_KEY = "mlflow.prompt.is_prompt"
# A special tag in ModelVersion to store the prompt text
PROMPT_TEXT_TAG_KEY = "mlflow.prompt.text"
# TODO: Replace this with model_ids in MLflow 3
PROMPT_ASSOCIATED_RUN_IDS_TAG_KEY = "mlflow.prompt.run_ids"
PROMPT_TEMPLATE_VARIABLE_PATTERN = re.compile(
r"\{\{\s*([a-zA-Z_][a-zA-Z0-9_]*(?:\.[a-zA-Z_][a-zA-Z0-9_]*)*)\s*\}\}"
)
PROMPT_TEXT_DISPLAY_LIMIT = 30
# Alphanumeric, underscore, hyphen, and dot are allowed in prompt name
PROMPT_NAME_RULE = re.compile(r"^[a-zA-Z0-9_.-]+$")

View File

@@ -0,0 +1,193 @@
import os
import re
import yaml
from mlflow.exceptions import MlflowException
from mlflow.version import VERSION as __version__
class _PromptlabModel:
import pandas as pd
def __init__(self, prompt_template, prompt_parameters, model_parameters, model_route):
self.prompt_parameters = prompt_parameters
self.model_parameters = model_parameters
self.model_route = model_route
self.prompt_template = prompt_template
def predict(self, inputs: pd.DataFrame) -> list[str]:
from mlflow.gateway import query
results = []
for idx in inputs.index:
prompt_parameters_as_dict = {
param.key: inputs[param.key][idx] for param in self.prompt_parameters
}
# copy replacement logic from PromptEngineering.utils.ts for consistency
prompt = self.prompt_template
for key, value in prompt_parameters_as_dict.items():
prompt = re.sub(r"\{\{\s*" + key + r"\s*\}\}", value, prompt)
model_parameters_as_dict = {param.key: param.value for param in self.model_parameters}
query_data = self._construct_query_data(prompt)
response = query(
route=self.model_route, data={**query_data, **model_parameters_as_dict}
)
results.append(self._parse_gateway_response(response))
return results
def _construct_query_data(self, prompt):
from mlflow.gateway import get_route
route_type = get_route(self.model_route).route_type
if route_type == "llm/v1/completions":
return {"prompt": prompt}
elif route_type == "llm/v1/chat":
return {"messages": [{"content": prompt, "role": "user"}]}
else:
raise MlflowException(
"Error when constructing gateway query: "
f"Unsupported route type for _PromptlabModel: {route_type}"
)
def _parse_gateway_response(self, response):
from mlflow.gateway import get_route
route_type = get_route(self.model_route).route_type
if route_type == "llm/v1/completions":
return response["choices"][0]["text"]
elif route_type == "llm/v1/chat":
return response["choices"][0]["message"]["content"]
else:
raise MlflowException(
"Error when parsing gateway response: "
f"Unsupported route type for _PromptlabModel: {route_type}"
)
def _load_pyfunc(path):
from mlflow import pyfunc
from mlflow.entities.param import Param
from mlflow.utils.model_utils import (
_get_flavor_configuration,
)
pyfunc_flavor_conf = _get_flavor_configuration(model_path=path, flavor_name=pyfunc.FLAVOR_NAME)
parameters_path = os.path.join(path, pyfunc_flavor_conf["parameters_path"])
with open(parameters_path) as f:
parameters = yaml.safe_load(f)
prompt_parameters_as_params = [
Param(key=key, value=value) for key, value in parameters["prompt_parameters"].items()
]
model_parameters_as_params = [
Param(key=key, value=value) for key, value in parameters["model_parameters"].items()
]
return _PromptlabModel(
prompt_template=parameters["prompt_template"],
prompt_parameters=prompt_parameters_as_params,
model_parameters=model_parameters_as_params,
model_route=parameters["model_route"],
)
def save_model(
path,
conda_env=None,
code_paths=None,
mlflow_model=None,
signature=None,
input_example=None,
pip_requirements=None,
prompt_template=None,
prompt_parameters=None,
model_parameters=None,
model_route=None,
):
from mlflow import pyfunc
from mlflow.models import Model
from mlflow.models.model import MLMODEL_FILE_NAME, Model
from mlflow.models.utils import _save_example
from mlflow.utils.environment import (
_CONDA_ENV_FILE_NAME,
_CONSTRAINTS_FILE_NAME,
_PYTHON_ENV_FILE_NAME,
_REQUIREMENTS_FILE_NAME,
_process_conda_env,
_process_pip_requirements,
_PythonEnv,
_validate_env_arguments,
infer_pip_requirements,
)
from mlflow.utils.file_utils import write_to
from mlflow.utils.model_utils import (
_validate_and_copy_code_paths,
_validate_and_prepare_target_save_path,
)
_validate_env_arguments(conda_env, pip_requirements, None)
_validate_and_prepare_target_save_path(path)
code_dir_subpath = _validate_and_copy_code_paths(code_paths, path)
if mlflow_model is None:
mlflow_model = Model()
if signature is not None:
mlflow_model.signature = signature
if input_example is not None:
_save_example(mlflow_model, input_example, path)
parameters_sub_path = "parameters.yaml"
parameters_path = os.path.join(path, parameters_sub_path)
# dump prompt_template, prompt_parameters, model_parameters, model_route to parameters_path
parameters = {
"prompt_template": prompt_template,
"prompt_parameters": {param.key: param.value for param in prompt_parameters},
"model_parameters": {param.key: param.value for param in model_parameters},
"model_route": model_route,
}
with open(parameters_path, "w") as f:
yaml.safe_dump(parameters, stream=f, default_flow_style=False)
pyfunc.add_to_model(
mlflow_model,
loader_module="mlflow.prompt.promptlab_model",
parameters_path=parameters_sub_path,
conda_env=_CONDA_ENV_FILE_NAME,
python_env=_PYTHON_ENV_FILE_NAME,
code=code_dir_subpath,
)
mlflow_model.save(os.path.join(path, MLMODEL_FILE_NAME))
if conda_env is None:
if pip_requirements is None:
inferred_reqs = infer_pip_requirements(
path, "mlflow._promptlab", [f"mlflow[gateway]=={__version__}"]
)
default_reqs = sorted(inferred_reqs)
else:
default_reqs = None
conda_env, pip_requirements, pip_constraints = _process_pip_requirements(
default_reqs, pip_requirements, None
)
else:
conda_env, pip_requirements, pip_constraints = _process_conda_env(conda_env)
with open(os.path.join(path, _CONDA_ENV_FILE_NAME), "w") as f:
yaml.safe_dump(conda_env, stream=f, default_flow_style=False)
if pip_constraints:
write_to(os.path.join(path, _CONSTRAINTS_FILE_NAME), "\n".join(pip_constraints))
write_to(os.path.join(path, _REQUIREMENTS_FILE_NAME), "\n".join(pip_requirements))
_PythonEnv.current().to_yaml(os.path.join(path, _PYTHON_ENV_FILE_NAME))

View File

@@ -0,0 +1,146 @@
import functools
import re
from textwrap import dedent
from typing import Any, Optional, Union
import mlflow
from mlflow.entities.model_registry.registered_model_tag import RegisteredModelTag
from mlflow.exceptions import MlflowException
from mlflow.prompt.constants import IS_PROMPT_TAG_KEY, PROMPT_NAME_RULE
from mlflow.protos.databricks_pb2 import RESOURCE_ALREADY_EXISTS
def add_prompt_filter_string(
filter_string: Optional[str], is_prompt: bool = False
) -> Optional[str]:
"""
Additional filter string to include/exclude prompts from the result.
By default, exclude prompts from the result.
"""
if IS_PROMPT_TAG_KEY not in (filter_string or ""):
prompt_filter_query = (
f"tag.`{IS_PROMPT_TAG_KEY}` = 'true'"
if is_prompt
else f"tag.`{IS_PROMPT_TAG_KEY}` != 'true'"
)
if filter_string:
filter_string = f"{filter_string} AND {prompt_filter_query}"
else:
filter_string = prompt_filter_query
return filter_string
def has_prompt_tag(tags: Optional[Union[list[RegisteredModelTag], dict[str, str]]]) -> bool:
"""Check if the given tags contain the prompt tag."""
if isinstance(tags, dict):
return IS_PROMPT_TAG_KEY in tags if tags else False
if not tags:
return
return any(tag.key == IS_PROMPT_TAG_KEY for tag in tags)
def is_prompt_supported_registry(registry_uri: Optional[str] = None) -> bool:
"""
Check if the current registry supports prompts.
Prompts registration is supported only in the OSS MLflow Tracking Server,
not in Databricks or OSS Unity Catalog.
"""
registry_uri = registry_uri or mlflow.get_registry_uri()
return not registry_uri.startswith("databricks") and not registry_uri.startswith("uc:")
def require_prompt_registry(func):
"""Ensure that the current registry supports prompts."""
@functools.wraps(func)
def wrapper(*args, **kwargs):
if args and isinstance(args[0], mlflow.MlflowClient):
registry_uri = args[0]._registry_uri
else:
registry_uri = mlflow.get_registry_uri()
if not is_prompt_supported_registry(registry_uri):
raise MlflowException(
f"The '{func.__name__}' API is only available with the OSS MLflow Tracking Server."
)
return func(*args, **kwargs)
# Add note about prompt support to the docstring
func.__doc__ = dedent(f"""\
{func.__doc__}
.. note::
This API is supported only when using the OSS MLflow Model Registry. Prompts are not
supported in Databricks or the OSS Unity Catalog model registry.
""")
return wrapper
def translate_prompt_exception(func):
"""
Translate MlflowException message related to RegisteredModel / ModelVersion into
prompt-specific message.
"""
MODEL_PATTERN = re.compile(r"(registered model|model version)", re.IGNORECASE)
@functools.wraps(func)
def wrapper(*args, **kwargs):
try:
return func(*args, **kwargs)
except MlflowException as e:
original_message = e.message
# Preserve the case of the first letter
new_message = MODEL_PATTERN.sub(
lambda m: "Prompt" if m.group(0)[0].isupper() else "prompt", e.message
)
if new_message != original_message:
raise MlflowException(new_message) from e
else:
raise e
return wrapper
def validate_prompt_name(name: Any):
"""Validate the prompt name against the prompt specific rule"""
if not isinstance(name, str) or not name:
raise MlflowException.invalid_parameter_value(
"Prompt name must be a non-empty string.",
)
if PROMPT_NAME_RULE.match(name) is None:
raise MlflowException.invalid_parameter_value(
"Prompt name can only contain alphanumeric characters, hyphens, underscores, and dots.",
)
def handle_resource_already_exist_error(
name: str,
is_existing_entity_prompt: bool,
is_new_entity_prompt: bool,
):
"""
Show a more specific error message for name conflict in Model Registry.
1. When creating a model with the same name as an existing model, say "model already exists".
2. When creating a prompt with the same name as an existing prompt, say "prompt already exists".
3. Otherwise, explain that a prompt and a model cannot have the same name.
"""
old_entity = "Prompt" if is_existing_entity_prompt else "Registered Model"
new_entity = "Prompt" if is_new_entity_prompt else "Registered Model"
if old_entity != new_entity:
raise MlflowException(
f"Tried to create a {new_entity.lower()} with name {name!r}, but the name is "
f"already taken by a {old_entity.lower()}. MLflow does not allow creating a "
"model and a prompt with the same name.",
RESOURCE_ALREADY_EXISTS,
)
raise MlflowException(
f"{new_entity} (name={name}) already exists.",
RESOURCE_ALREADY_EXISTS,
)