Files
zenml/venv/lib/python3.9/site-packages/mlflow/gateway/app.py
Christian Mantha 2ca0b9ef7c star
2026-03-02 19:10:52 -05:00

437 lines
16 KiB
Python

import functools
from pathlib import Path
from typing import Any, Optional, Union
from fastapi import FastAPI, HTTPException, Request
from fastapi.openapi.docs import get_swagger_ui_html
from fastapi.responses import FileResponse, RedirectResponse
from pydantic import BaseModel
from slowapi import Limiter, _rate_limit_exceeded_handler
from slowapi.errors import RateLimitExceeded
from slowapi.util import get_remote_address
from mlflow.deployments.server.config import Endpoint
from mlflow.deployments.server.constants import (
MLFLOW_DEPLOYMENTS_CRUD_ENDPOINT_BASE,
MLFLOW_DEPLOYMENTS_ENDPOINTS_BASE,
MLFLOW_DEPLOYMENTS_HEALTH_ENDPOINT,
MLFLOW_DEPLOYMENTS_LIMITS_BASE,
MLFLOW_DEPLOYMENTS_LIST_ENDPOINTS_PAGE_SIZE,
MLFLOW_DEPLOYMENTS_QUERY_SUFFIX,
)
from mlflow.environment_variables import (
MLFLOW_GATEWAY_CONFIG,
MLFLOW_GATEWAY_RATE_LIMITS_STORAGE_URI,
)
from mlflow.exceptions import MlflowException
from mlflow.gateway.base_models import SetLimitsModel
from mlflow.gateway.config import (
GatewayConfig,
LimitsConfig,
Route,
RouteConfig,
RouteType,
_load_route_config,
)
from mlflow.gateway.constants import (
MLFLOW_GATEWAY_CRUD_ROUTE_BASE,
MLFLOW_GATEWAY_HEALTH_ENDPOINT,
MLFLOW_GATEWAY_LIMITS_BASE,
MLFLOW_GATEWAY_ROUTE_BASE,
MLFLOW_GATEWAY_SEARCH_ROUTES_PAGE_SIZE,
MLFLOW_QUERY_SUFFIX,
)
from mlflow.gateway.exceptions import AIGatewayException
from mlflow.gateway.providers import get_provider
from mlflow.gateway.schemas import chat, completions, embeddings
from mlflow.gateway.utils import SearchRoutesToken, make_streaming_response
from mlflow.version import VERSION
class GatewayAPI(FastAPI):
def __init__(self, config: GatewayConfig, limiter: Limiter, *args: Any, **kwargs: Any):
super().__init__(*args, **kwargs)
self.state.limiter = limiter
self.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
self.dynamic_routes: dict[str, RouteConfig] = {}
self.set_dynamic_routes(config, limiter)
def set_dynamic_routes(self, config: GatewayConfig, limiter: Limiter) -> None:
self.dynamic_routes.clear()
for route in config.routes:
# TODO: Remove deployments server URLs after deprecation window elapses
self.add_api_route(
path=(
MLFLOW_DEPLOYMENTS_ENDPOINTS_BASE + route.name + MLFLOW_DEPLOYMENTS_QUERY_SUFFIX
),
endpoint=_route_type_to_endpoint(route, limiter, "deployments"),
methods=["POST"],
)
self.add_api_route(
path=f"{MLFLOW_GATEWAY_ROUTE_BASE}{route.name}{MLFLOW_QUERY_SUFFIX}",
endpoint=_route_type_to_endpoint(route, limiter, "gateway"),
methods=["POST"],
include_in_schema=False,
)
self.dynamic_routes[route.name] = route
def get_dynamic_route(self, route_name: str) -> Optional[Route]:
return r.to_route() if (r := self.dynamic_routes.get(route_name)) else None
def _translate_http_exception(func):
"""
Decorator for translating MLflow exceptions to HTTP exceptions
"""
@functools.wraps(func)
async def wrapper(*args, **kwargs):
try:
return await func(*args, **kwargs)
except AIGatewayException as e:
raise HTTPException(status_code=e.status_code, detail=e.detail)
return wrapper
def _create_chat_endpoint(config: RouteConfig):
prov = get_provider(config.model.provider)(config)
# https://slowapi.readthedocs.io/en/latest/#limitations-and-known-issues
@_translate_http_exception
async def _chat(
request: Request, payload: chat.RequestPayload
) -> Union[chat.ResponsePayload, chat.StreamResponsePayload]:
if payload.stream:
return await make_streaming_response(prov.chat_stream(payload))
else:
return await prov.chat(payload)
return _chat
def _create_completions_endpoint(config: RouteConfig):
prov = get_provider(config.model.provider)(config)
@_translate_http_exception
async def _completions(
request: Request, payload: completions.RequestPayload
) -> Union[completions.ResponsePayload, completions.StreamResponsePayload]:
if payload.stream:
return await make_streaming_response(prov.completions_stream(payload))
else:
return await prov.completions(payload)
return _completions
def _create_embeddings_endpoint(config: RouteConfig):
prov = get_provider(config.model.provider)(config)
@_translate_http_exception
async def _embeddings(
request: Request, payload: embeddings.RequestPayload
) -> embeddings.ResponsePayload:
return await prov.embeddings(payload)
return _embeddings
async def _custom(request: Request):
return request.json()
def _route_type_to_endpoint(config: RouteConfig, limiter: Limiter, key: str):
provider_to_factory = {
RouteType.LLM_V1_CHAT: _create_chat_endpoint,
RouteType.LLM_V1_COMPLETIONS: _create_completions_endpoint,
RouteType.LLM_V1_EMBEDDINGS: _create_embeddings_endpoint,
}
if factory := provider_to_factory.get(config.route_type):
handler = factory(config)
if limit := config.limit:
limit_value = f"{limit.calls}/{limit.renewal_period}"
handler.__name__ = f"{handler.__name__}_{config.name}_{key}"
return limiter.limit(limit_value)(handler)
else:
return handler
raise HTTPException(
status_code=404,
detail=f"Unexpected route type {config.route_type!r} for route {config.name!r}.",
)
class HealthResponse(BaseModel):
status: str
class ListEndpointsResponse(BaseModel):
endpoints: list[Endpoint]
next_page_token: Optional[str] = None
class Config:
schema_extra = {
"example": {
"endpoints": [
{
"name": "openai-chat",
"endpoint_type": "llm/v1/chat",
"model": {
"name": "gpt-4o-mini",
"provider": "openai",
},
"limit": {"calls": 1, "key": None, "renewal_period": "minute"},
},
{
"name": "anthropic-completions",
"endpoint_type": "llm/v1/completions",
"model": {
"name": "claude-instant-100k",
"provider": "anthropic",
},
},
{
"name": "cohere-embeddings",
"endpoint_type": "llm/v1/embeddings",
"model": {
"name": "embed-english-v2.0",
"provider": "cohere",
},
},
],
"next_page_token": "eyJpbmRleCI6IDExfQ==",
}
}
class SearchRoutesResponse(BaseModel):
routes: list[Route]
next_page_token: Optional[str] = None
class Config:
schema_extra = {
"example": {
"routes": [
{
"name": "openai-chat",
"route_type": "llm/v1/chat",
"model": {
"name": "gpt-4o-mini",
"provider": "openai",
},
},
{
"name": "anthropic-completions",
"route_type": "llm/v1/completions",
"model": {
"name": "claude-instant-100k",
"provider": "anthropic",
},
},
{
"name": "cohere-embeddings",
"route_type": "llm/v1/embeddings",
"model": {
"name": "embed-english-v2.0",
"provider": "cohere",
},
},
],
"next_page_token": "eyJpbmRleCI6IDExfQ==",
}
}
def create_app_from_config(config: GatewayConfig) -> GatewayAPI:
"""
Create the GatewayAPI app from the gateway configuration.
"""
limiter = Limiter(
key_func=get_remote_address, storage_uri=MLFLOW_GATEWAY_RATE_LIMITS_STORAGE_URI.get()
)
app = GatewayAPI(
config=config,
limiter=limiter,
title="MLflow AI Gateway",
description="The core deployments API for reverse proxy interface using remote inference "
"endpoints within MLflow",
version=VERSION,
docs_url=None,
)
@app.get("/", include_in_schema=False)
async def index():
return RedirectResponse(url="/docs")
@app.get("/favicon.ico", include_in_schema=False)
async def favicon():
for directory in ["build", "public"]:
favicon_file = Path(__file__).parent.parent.joinpath(
"server", "js", directory, "favicon.ico"
)
if favicon_file.exists():
return FileResponse(favicon_file)
raise HTTPException(status_code=404, detail="favicon.ico not found")
@app.get("/docs", include_in_schema=False)
async def docs():
return get_swagger_ui_html(
openapi_url="/openapi.json",
title="MLflow AI Gateway",
swagger_favicon_url="/favicon.ico",
)
# TODO: Remove deployments server URLs after deprecation window elapses
@app.get(MLFLOW_DEPLOYMENTS_HEALTH_ENDPOINT)
@app.get(MLFLOW_GATEWAY_HEALTH_ENDPOINT, include_in_schema=False)
async def health() -> HealthResponse:
return {"status": "OK"}
# TODO: Remove deployments server URLs after deprecation window elapses
@app.get(MLFLOW_DEPLOYMENTS_CRUD_ENDPOINT_BASE + "{endpoint_name}")
async def get_endpoint(endpoint_name: str) -> Endpoint:
if matched := app.get_dynamic_route(endpoint_name):
return matched.to_endpoint()
raise HTTPException(
status_code=404,
detail=f"The endpoint '{endpoint_name}' is not present or active on the server. Please "
"verify the endpoint name.",
)
@app.get(MLFLOW_GATEWAY_CRUD_ROUTE_BASE + "{route_name}", include_in_schema=False)
async def get_route(route_name: str) -> Route:
if matched := app.get_dynamic_route(route_name):
return matched
raise HTTPException(
status_code=404,
detail=f"The route '{route_name}' is not present or active on the server. Please "
"verify the route name.",
)
# TODO: Remove deployments server URLs after deprecation window elapses
@app.get(MLFLOW_DEPLOYMENTS_CRUD_ENDPOINT_BASE)
async def list_endpoints(page_token: Optional[str] = None) -> ListEndpointsResponse:
start_idx = SearchRoutesToken.decode(page_token).index if page_token is not None else 0
end_idx = start_idx + MLFLOW_DEPLOYMENTS_LIST_ENDPOINTS_PAGE_SIZE
routes = list(app.dynamic_routes.values())
result = {
"endpoints": [route.to_route().to_endpoint() for route in routes[start_idx:end_idx]]
}
if len(routes[end_idx:]) > 0:
next_page_token = SearchRoutesToken(index=end_idx)
result["next_page_token"] = next_page_token.encode()
return result
@app.get(MLFLOW_GATEWAY_CRUD_ROUTE_BASE, include_in_schema=False)
async def search_routes(page_token: Optional[str] = None) -> SearchRoutesResponse:
start_idx = SearchRoutesToken.decode(page_token).index if page_token is not None else 0
end_idx = start_idx + MLFLOW_GATEWAY_SEARCH_ROUTES_PAGE_SIZE
routes = list(app.dynamic_routes.values())
result = {"routes": [r.to_route() for r in routes[start_idx:end_idx]]}
if len(routes[end_idx:]) > 0:
next_page_token = SearchRoutesToken(index=end_idx)
result["next_page_token"] = next_page_token.encode()
return result
# TODO: Remove deployments server URLs after deprecation window elapses
@app.get(MLFLOW_DEPLOYMENTS_LIMITS_BASE + "{endpoint}")
@app.get(MLFLOW_GATEWAY_LIMITS_BASE + "{endpoint}", include_in_schema=False)
async def get_limits(endpoint: str) -> LimitsConfig:
raise HTTPException(status_code=501, detail="The get_limits API is not available yet.")
# TODO: Remove deployments server URLs after deprecation window elapses
@app.post(MLFLOW_DEPLOYMENTS_LIMITS_BASE)
@app.post(MLFLOW_GATEWAY_LIMITS_BASE, include_in_schema=False)
async def set_limits(payload: SetLimitsModel) -> LimitsConfig:
raise HTTPException(status_code=501, detail="The set_limits API is not available yet.")
def _look_up_route(name: str) -> Optional[Route]:
if r := app.dynamic_routes.get(name):
return r
raise HTTPException(
status_code=400,
detail=f"Route {name} not found in the configuration.",
)
@app.post("/v1/chat/completions")
async def openai_chat_handler(
request: Request, payload: chat.RequestPayload
) -> chat.ResponsePayload:
route = _look_up_route(payload.model)
if route.route_type != RouteType.LLM_V1_CHAT:
raise HTTPException(
status_code=400,
detail=f"Endpoint {route.name!r} is not a chat endpoint.",
)
prov = get_provider(route.model.provider)(route)
payload.model = None # provider rejects a request with model field, must be set to None
if payload.stream:
return await make_streaming_response(prov.chat_stream(payload))
else:
return await prov.chat(payload)
@app.post("/v1/completions")
async def openai_completions_handler(
request: Request, payload: completions.RequestPayload
) -> completions.ResponsePayload:
route = _look_up_route(payload.model)
if route.route_type != RouteType.LLM_V1_COMPLETIONS:
raise HTTPException(
status_code=400,
detail=f"Endpoint {route.name!r} is not a completions endpoint.",
)
prov = get_provider(route.model.provider)(route)
payload.model = None # provider rejects a request with model field, must be set to None
if payload.stream:
return await make_streaming_response(prov.completions_stream(payload))
else:
return await prov.completions(payload)
@app.post("/v1/embeddings")
async def openai_embeddings_handler(
request: Request, payload: embeddings.RequestPayload
) -> embeddings.ResponsePayload:
route = _look_up_route(payload.model)
if route.route_type != RouteType.LLM_V1_EMBEDDINGS:
raise HTTPException(
status_code=400,
detail=f"Endpoint {route.name!r} is not an embeddings endpoint.",
)
prov = get_provider(route.model.provider)(route)
payload.model = None # provider rejects a request with model field, must be set to None
return await prov.embeddings(payload)
return app
def create_app_from_path(config_path: Union[str, Path]) -> GatewayAPI:
"""
Load the path and generate the GatewayAPI app instance.
"""
config = _load_route_config(config_path)
return create_app_from_config(config)
def create_app_from_env() -> GatewayAPI:
"""
Load the path from the environment variable and generate the GatewayAPI app instance.
"""
if config_path := MLFLOW_GATEWAY_CONFIG.get():
return create_app_from_path(config_path)
raise MlflowException(
f"Environment variable {MLFLOW_GATEWAY_CONFIG!r} is not set. "
"Please set it to the path of the gateway configuration file."
)