diff --git a/manifests/helm/mlflow/1.9.0/CUSTOM-README.md b/manifests/helm/mlflow/1.9.0/CUSTOM-README.md index 1835f34..1408edf 100644 --- a/manifests/helm/mlflow/1.9.0/CUSTOM-README.md +++ b/manifests/helm/mlflow/1.9.0/CUSTOM-README.md @@ -231,6 +231,7 @@ extraArgs: appName: "oidc-auth" uvicornOpts: "--timeout-keep-alive 600" allowedHosts: "mlflow.example.org" + corsAllowedOrigins: "https://mlflow.example.org" # CORS 허용 Origin log: enabled: false # uvicornOpts 사용 시 반드시 false (gunicorn/uvicorn 충돌 방지) @@ -272,26 +273,46 @@ extraVolumeMounts: --- -### 3.7 OIDC Auth Middleware 패치 +### 3.7 before_request.py 패치 -`mlflow-oidc-auth` 플러그인의 `auth_middleware.py`를 차트에 포함된 버전으로 교체한다. -워크스페이스 지원(`x-mlflow-workspace` 헤더 처리) 등 업스트림 수정 사항을 반영한다. +`mlflow-oidc-auth` v7.0.3의 `hooks/before_request.py` line 520에 권한 체크 오류가 있다. +workspace에서 실험·모델 생성 시 MANAGE(`can_manage`)를 요구하지만, +`OIDC_WORKSPACE_DEFAULT_PERMISSION: "EDIT"`으로 자동 부여된 EDIT 권한은 +`can_update=True, can_manage=False`이므로 EDIT 사용자가 항상 403을 받는다. -```yaml -oidcAuthPatch: - enabled: true - mountPath: "/usr/local/lib/python3.11/site-packages/mlflow_oidc_auth/middleware/auth_middleware.py" +#### 수정 내용 (`files/before_request.py` line 520) + +```python +# 원본 (버그) +if ws_perm is None or not ws_perm.can_manage: + return responses.make_forbidden_response() + +# 패치 (수정) +if ws_perm is None or not ws_perm.can_update: + return responses.make_forbidden_response() ``` -파일 소스: `files/auth_middleware.py` +#### ConfigMap 생성 -> **Python 버전 확인**: 컨테이너 이미지의 Python 버전이 다를 경우 `mountPath`를 수정한다. +```sh +kubectl create configmap mlflow-hooks-patch -n mlflow \ + --from-file=before_request.py=manifests/helm/mlflow/1.9.0/files/before_request.py +``` + +패치 파일 소스: `files/before_request.py` + +> **Python 버전 확인**: 컨테이너 이미지의 Python 버전이 다를 경우 `extraVolumeMounts`의 `mountPath`를 수정한다. +> `paasup/mlflow:v3.11.1-oidc` 이미지의 실제 Python 버전은 **3.10**이므로 경로는 `python3.10`을 사용한다. > > ```sh > kubectl exec -n mlflow -- python -c \ -> "import mlflow_oidc_auth.middleware.auth_middleware as m; print(m.__file__)" +> "import mlflow_oidc_auth.hooks.before_request as m; print(m.__file__)" > ``` +> **업스트림 이슈**: 이 동작이 의도된 설계인지 여부를 mlflow-oidc-auth 저장소에 문의했다. +> → [Issue #240](https://github.com/mlflow-oidc/mlflow-oidc-auth/issues/240) +> 업스트림에서 수정이 반영되면 이 패치는 제거한다. + --- ### 3.8 Ingress 설정 @@ -319,7 +340,87 @@ ingress: --- -### 3.9 PostgreSQL 설정 +### 3.9 CORS 설정 + +브라우저가 MLflow API를 cross-origin으로 호출할 때 발생하는 `Cross-origin request blocked`를 해결하기 위해 두 가지 설정이 필요하다. + +#### MLflow 서버 네이티브 설정 + +`extraArgs.corsAllowedOrigins`으로 `--cors-allowed-origins` 플래그를 전달한다. + +```yaml +extraArgs: + corsAllowedOrigins: "https://mlflow.example.org" +``` + +여러 origin을 허용할 경우 쉼표로 구분한다. + +```yaml + corsAllowedOrigins: "https://mlflow.example.org,https://other.example.org" +``` + +#### Kong Ingress CORS 플러그인 + +Kong이 CORS preflight(OPTIONS) 요청을 처리하고 응답 헤더를 보완하도록 KongPlugin을 추가한다. + +```sh +kubectl apply -f - < **주의**: 브라우저 콘솔에서 403 응답이 `Cross-origin request blocked`로 표시되는 경우가 있다. +> 응답 헤더에 `Access-Control-Allow-Origin`이 존재하면 CORS 자체는 정상이며, 실제 원인은 MLflow 권한(403) 문제이다. + +--- + +### 3.10 PostgreSQL 설정 내장 PostgreSQL을 사용한다. @@ -346,7 +447,7 @@ extraEnvVars: --- -### 3.10 S3 (MinIO) 설정 +### 3.11 S3 (MinIO) 설정 ```yaml artifactRoot: diff --git a/manifests/helm/mlflow/1.9.0/custom-values.yaml b/manifests/helm/mlflow/1.9.0/custom-values.yaml index b1c61b7..0d96c96 100644 --- a/manifests/helm/mlflow/1.9.0/custom-values.yaml +++ b/manifests/helm/mlflow/1.9.0/custom-values.yaml @@ -68,6 +68,7 @@ extraEnvVars: WORKSPACE_CACHE_MAX_SIZE: "1024" WORKSPACE_CACHE_TTL_SECONDS: "300" PERMISSION_SOURCE_ORDER: "user,group,regex,group-regex" + MLFLOW_LOGGING_LEVEL: "DEBUG" extraSecretNamesForEnvFrom: - mlflow-oidc-secret @@ -82,6 +83,7 @@ extraArgs: appName: "oidc-auth" uvicornOpts: "--timeout-keep-alive 600" allowedHosts: "mlflow.example.org" + corsAllowedOrigins: "https://mlflow.example.org" service: type: ClusterIP @@ -94,6 +96,7 @@ ingress: cert-manager.io/cluster-issuer: "selfsigned-issuer" cert-manager.io/duration: 8760h cert-manager.io/renew-before: 720h + konghq.com/plugins: mlflow-cors hosts: - host: mlflow.example.org paths: @@ -119,6 +122,9 @@ extraVolumes: - name: workspace-plugin configMap: name: mlflow-workspace-plugin + - name: hooks-patch + configMap: + name: mlflow-hooks-patch extraVolumeMounts: - name: keycloak-ca-cert @@ -127,6 +133,10 @@ extraVolumeMounts: readOnly: true - name: workspace-plugin mountPath: /opt/mlflow-plugins + - name: hooks-patch + mountPath: /usr/local/lib/python3.10/site-packages/mlflow_oidc_auth/hooks/before_request.py + subPath: before_request.py + readOnly: true serviceMonitor: enabled: false diff --git a/manifests/helm/mlflow/1.9.0/files/before_request.py b/manifests/helm/mlflow/1.9.0/files/before_request.py new file mode 100644 index 0000000..a39ef21 --- /dev/null +++ b/manifests/helm/mlflow/1.9.0/files/before_request.py @@ -0,0 +1,598 @@ +import re +from typing import Any, Callable, Dict, Optional + +from flask import Request, g, request +from mlflow.protos.model_registry_pb2 import ( + CreateModelVersion, + DeleteModelVersion, + DeleteModelVersionTag, + DeleteRegisteredModel, + DeleteRegisteredModelAlias, + DeleteRegisteredModelTag, + GetLatestVersions, + GetModelVersion, + GetModelVersionByAlias, + GetModelVersionDownloadUri, + GetRegisteredModel, + RenameRegisteredModel, + SetModelVersionTag, + SetRegisteredModelAlias, + SetRegisteredModelTag, + TransitionModelVersionStage, + UpdateModelVersion, + UpdateRegisteredModel, +) +from mlflow.protos.service_pb2 import ( + AttachModelToGatewayEndpoint, + CreateGatewayEndpoint, + CreateGatewayEndpointBinding, + CreateGatewayModelDefinition, + CreateGatewaySecret, + CreateLoggedModel, + CreateRun, + CreateWorkspace, + DeleteExperiment, + DeleteExperimentTag, + DeleteGatewayEndpoint, + DeleteGatewayEndpointBinding, + DeleteGatewayEndpointTag, + DeleteGatewayModelDefinition, + DeleteGatewaySecret, + DeleteLoggedModel, + DeleteLoggedModelTag, + DeleteRun, + DeleteTag, + DeleteWorkspace, + DetachModelFromGatewayEndpoint, + FinalizeLoggedModel, + GetExperiment, + GetExperimentByName, + GetGatewayEndpoint, + GetGatewayModelDefinition, + GetGatewaySecretInfo, + GetLoggedModel, + GetMetricHistory, + GetRun, + GetWorkspace, + ListArtifacts, + ListGatewayEndpointBindings, + ListWorkspaces, + LogBatch, + LogLoggedModelParamsRequest, + LogMetric, + LogModel, + LogParam, + RestoreExperiment, + RestoreRun, + SetExperimentTag, + SetGatewayEndpointTag, + SetLoggedModelTags, + SetTag, + UpdateExperiment, + UpdateGatewayEndpoint, + UpdateGatewayModelDefinition, + UpdateGatewaySecret, + UpdateRun, + UpdateWorkspace, + RegisterScorer, + ListScorers, + GetScorer, + DeleteScorer, + ListScorerVersions, + CreatePromptOptimizationJob, + GetPromptOptimizationJob, + SearchPromptOptimizationJobs, + DeletePromptOptimizationJob, + CancelPromptOptimizationJob, +) + +from mlflow.server.handlers import catch_mlflow_exception, get_endpoints +from mlflow.utils.rest_utils import _REST_API_PATH_PREFIX + +# Forward-compatible imports for Gateway Budget Policy protos. +# These protos may not exist in the installed MLflow version; when they +# become available they will be automatically picked up as admin-only handlers. +_BUDGET_POLICY_PROTOS: list = [] +try: + from mlflow.protos.service_pb2 import ( + CreateGatewayBudgetPolicy, + UpdateGatewayBudgetPolicy, + DeleteGatewayBudgetPolicy, + ) + + _BUDGET_POLICY_PROTOS = [ + CreateGatewayBudgetPolicy, + UpdateGatewayBudgetPolicy, + DeleteGatewayBudgetPolicy, + ] +except ImportError: + pass + +from mlflow_oidc_auth.bridge import get_fastapi_admin_status, get_fastapi_username +import mlflow_oidc_auth.responses as responses +from mlflow_oidc_auth.config import config +from mlflow_oidc_auth.logger import get_logger +from mlflow_oidc_auth.validators import ( + validate_can_delete_experiment, + validate_can_delete_experiment_artifact_proxy, + validate_can_delete_logged_model, + validate_can_delete_registered_model, + validate_can_delete_run, + validate_can_manage_experiment, + validate_can_manage_registered_model, + validate_can_read_experiment, + validate_can_read_experiment_artifact_proxy, + validate_can_read_experiment_by_name, + validate_can_read_logged_model, + validate_can_read_registered_model, + validate_can_read_run, + validate_can_update_experiment, + validate_can_update_experiment_artifact_proxy, + validate_can_update_logged_model, + validate_can_update_registered_model, + validate_can_update_run, + validate_can_read_experiments_from_experiment_ids, + validate_can_update_experiment_from_experiment_id, + validate_can_read_metric_history_bulk_interval, + validate_can_read_traces_from_experiment_ids, + validate_can_read_trace, + validate_can_update_trace_from_experiment_id, + validate_can_update_trace_from_run_id, + validate_can_update_trace, + validate_can_delete_traces_from_experiment_id, + validate_can_delete_scorer, + validate_can_manage_scorer, + validate_can_manage_scorer_permission, + validate_can_read_scorer, + validate_can_update_scorer, + validate_can_read_run_artifact, + validate_can_update_run_artifact, + validate_can_read_model_version_artifact, + validate_can_read_trace_artifact, + validate_can_read_metric_history_bulk, + validate_can_search_datasets, + validate_can_create_promptlab_run, + validate_gateway_proxy, + validate_can_read_gateway_endpoint, + validate_can_update_gateway_endpoint, + validate_can_delete_gateway_endpoint, + validate_can_read_gateway_secret, + validate_can_update_gateway_secret, + validate_can_delete_gateway_secret, + validate_can_read_gateway_model_definition, + validate_can_update_gateway_model_definition, + validate_can_delete_gateway_model_definition, + validate_can_create_gateway, + validate_can_create_workspace, + validate_can_read_workspace, + validate_can_update_workspace, + validate_can_delete_workspace, + validate_can_list_workspaces, + validate_can_read_prompt_optimization_job, + validate_can_update_prompt_optimization_job, + validate_can_delete_prompt_optimization_job, +) + + +def _is_unprotected_route(path: str) -> bool: + return path.startswith( + ( + "/static", + "/favicon.ico", + "/health", + "/metrics", + "/docs", + "/redoc", + "/openapi.json", + ) + ) + + +def _deny_non_admin(_username: str) -> bool: + """Sentinel validator that always denies non-admin users. + + Admin users are short-circuited before validators run in before_request_hook, + so this function is only called for non-admin users and must always return False. + """ + return False + + +def _get_auth_context() -> tuple[Optional[str], bool]: + """Best-effort retrieval of auth context injected by FastAPI.""" + try: + username = get_fastapi_username() + except Exception: + username = None + + try: + is_admin = get_fastapi_admin_status() + except Exception: + is_admin = False + + return username, is_admin + + +BEFORE_REQUEST_HANDLERS = { + # Routes for experiments + GetExperiment: validate_can_read_experiment, + GetExperimentByName: validate_can_read_experiment_by_name, + DeleteExperiment: validate_can_delete_experiment, + RestoreExperiment: validate_can_delete_experiment, + UpdateExperiment: validate_can_update_experiment, + SetExperimentTag: validate_can_update_experiment, + DeleteExperimentTag: validate_can_update_experiment, + # Routes for runs + CreateRun: validate_can_update_experiment, + GetRun: validate_can_read_run, + DeleteRun: validate_can_delete_run, + RestoreRun: validate_can_delete_run, + UpdateRun: validate_can_update_run, + LogMetric: validate_can_update_run, + LogBatch: validate_can_update_run, + LogModel: validate_can_update_run, + SetTag: validate_can_update_run, + DeleteTag: validate_can_update_run, + LogParam: validate_can_update_run, + GetMetricHistory: validate_can_read_run, + ListArtifacts: validate_can_read_run, + # Routes for model registry + GetRegisteredModel: validate_can_read_registered_model, + DeleteRegisteredModel: validate_can_delete_registered_model, + UpdateRegisteredModel: validate_can_update_registered_model, + RenameRegisteredModel: validate_can_update_registered_model, + GetLatestVersions: validate_can_read_registered_model, + CreateModelVersion: validate_can_update_registered_model, + GetModelVersion: validate_can_read_registered_model, + DeleteModelVersion: validate_can_delete_registered_model, + UpdateModelVersion: validate_can_update_registered_model, + TransitionModelVersionStage: validate_can_update_registered_model, + GetModelVersionDownloadUri: validate_can_read_registered_model, + SetRegisteredModelTag: validate_can_update_registered_model, + DeleteRegisteredModelTag: validate_can_update_registered_model, + SetModelVersionTag: validate_can_update_registered_model, + DeleteModelVersionTag: validate_can_delete_registered_model, + SetRegisteredModelAlias: validate_can_update_registered_model, + DeleteRegisteredModelAlias: validate_can_delete_registered_model, + GetModelVersionByAlias: validate_can_read_registered_model, + # Routes for scorers + RegisterScorer: validate_can_update_experiment, + ListScorers: validate_can_read_experiment, + GetScorer: validate_can_read_scorer, + DeleteScorer: validate_can_delete_scorer, + ListScorerVersions: validate_can_read_scorer, + # Routes for prompt optimization jobs (resolved via job_id → experiment_id) + CreatePromptOptimizationJob: validate_can_update_experiment, + GetPromptOptimizationJob: validate_can_read_prompt_optimization_job, + SearchPromptOptimizationJobs: validate_can_read_experiment, + DeletePromptOptimizationJob: validate_can_delete_prompt_optimization_job, + CancelPromptOptimizationJob: validate_can_update_prompt_optimization_job, + # Routes for gateway endpoints + CreateGatewayEndpoint: validate_can_create_gateway, + GetGatewayEndpoint: validate_can_read_gateway_endpoint, + UpdateGatewayEndpoint: validate_can_update_gateway_endpoint, + DeleteGatewayEndpoint: validate_can_delete_gateway_endpoint, + # Routes for gateway secrets + CreateGatewaySecret: validate_can_create_gateway, + GetGatewaySecretInfo: validate_can_read_gateway_secret, + UpdateGatewaySecret: validate_can_update_gateway_secret, + DeleteGatewaySecret: validate_can_delete_gateway_secret, + # Routes for gateway model definitions + CreateGatewayModelDefinition: validate_can_create_gateway, + GetGatewayModelDefinition: validate_can_read_gateway_model_definition, + UpdateGatewayModelDefinition: validate_can_update_gateway_model_definition, + DeleteGatewayModelDefinition: validate_can_delete_gateway_model_definition, + # Routes for gateway endpoint-model mappings + AttachModelToGatewayEndpoint: validate_can_update_gateway_endpoint, + DetachModelFromGatewayEndpoint: validate_can_update_gateway_endpoint, + # Routes for gateway endpoint bindings + CreateGatewayEndpointBinding: validate_can_update_gateway_endpoint, + DeleteGatewayEndpointBinding: validate_can_update_gateway_endpoint, + ListGatewayEndpointBindings: validate_can_read_gateway_endpoint, + # Routes for gateway endpoint tags + SetGatewayEndpointTag: validate_can_update_gateway_endpoint, + DeleteGatewayEndpointTag: validate_can_update_gateway_endpoint, +} + +# Gateway Budget Policy protos are admin-only. They are conditionally +# available (forward-compat), so we add them after the dict is defined. +for _bp in _BUDGET_POLICY_PROTOS: + BEFORE_REQUEST_HANDLERS[_bp] = _deny_non_admin + +# `mlflow.server.handlers.get_endpoints()` also includes non-protobuf endpoints like `/graphql` +# and Gateway discovery routes, whose handlers are *not* our auth validators. We must not treat +# those as validators (they don't accept `username`), otherwise the hook will crash at runtime. +_PROTO_VALIDATORS = set(BEFORE_REQUEST_HANDLERS.values()) + + +logger = get_logger() + + +def _get_before_request_handler(request_class): + return BEFORE_REQUEST_HANDLERS.get(request_class) + + +BEFORE_REQUEST_VALIDATORS = { + (http_path, method): handler + for http_path, handler, methods in get_endpoints(_get_before_request_handler) + for method in methods + if handler in _PROTO_VALIDATORS +} + +from mlflow.server.handlers import _add_static_prefix, _get_ajax_path + +# Flask routes (not part of Protobuf API) +GET_ARTIFACT = _add_static_prefix("/get-artifact") +UPLOAD_ARTIFACT = _get_ajax_path("/mlflow/upload-artifact") +GET_MODEL_VERSION_ARTIFACT = _add_static_prefix("/model-versions/get-artifact") +GET_TRACE_ARTIFACT = _get_ajax_path("/mlflow/get-trace-artifact") +GET_METRIC_HISTORY_BULK = _get_ajax_path("/mlflow/metrics/get-history-bulk") +GET_METRIC_HISTORY_BULK_INTERVAL = _get_ajax_path("/mlflow/metrics/get-history-bulk-interval") +SEARCH_DATASETS = _get_ajax_path("/mlflow/experiments/search-datasets") +CREATE_PROMPTLAB_RUN = _get_ajax_path("/mlflow/runs/create-promptlab-run") +GATEWAY_PROXY = _get_ajax_path("/mlflow/gateway-proxy") +INVOKE_SCORER = _get_ajax_path("/mlflow/invocations/scorer") +GATEWAY_SUPPORTED_PROVIDERS = _get_ajax_path("/mlflow/gateway/supported-providers") +GATEWAY_SUPPORTED_MODELS = _get_ajax_path("/mlflow/gateway/supported-models") +GATEWAY_PROVIDER_CONFIG = _get_ajax_path("/mlflow/gateway/provider-config") +GATEWAY_SECRETS_CONFIG = _get_ajax_path("/mlflow/gateway/secrets-config") + + +# Flask routes (no proto mapping) +BEFORE_REQUEST_VALIDATORS.update( + { + (GET_ARTIFACT, "GET"): validate_can_read_run_artifact, + (UPLOAD_ARTIFACT, "POST"): validate_can_update_run_artifact, + (GET_MODEL_VERSION_ARTIFACT, "GET"): validate_can_read_model_version_artifact, + (GET_TRACE_ARTIFACT, "GET"): validate_can_read_trace_artifact, + (GET_METRIC_HISTORY_BULK, "GET"): validate_can_read_metric_history_bulk, + ( + GET_METRIC_HISTORY_BULK_INTERVAL, + "GET", + ): validate_can_read_metric_history_bulk_interval, + (SEARCH_DATASETS, "POST"): validate_can_search_datasets, + (CREATE_PROMPTLAB_RUN, "POST"): validate_can_create_promptlab_run, + (GATEWAY_PROXY, "GET"): validate_gateway_proxy, + (GATEWAY_PROXY, "POST"): validate_gateway_proxy, + # Scorer invocation uses the same gateway proxy permission check + (INVOKE_SCORER, "GET"): validate_gateway_proxy, + (INVOKE_SCORER, "POST"): validate_gateway_proxy, + # Gateway discovery routes use the same gateway proxy permission check + (GATEWAY_SUPPORTED_PROVIDERS, "GET"): validate_gateway_proxy, + (GATEWAY_SUPPORTED_MODELS, "GET"): validate_gateway_proxy, + # Gateway configuration routes are admin-only + (GATEWAY_PROVIDER_CONFIG, "GET"): _deny_non_admin, + (GATEWAY_SECRETS_CONFIG, "GET"): _deny_non_admin, + } +) + +LOGGED_MODEL_BEFORE_REQUEST_HANDLERS = { + CreateLoggedModel: validate_can_update_experiment, + GetLoggedModel: validate_can_read_logged_model, + DeleteLoggedModel: validate_can_delete_logged_model, + FinalizeLoggedModel: validate_can_update_logged_model, + DeleteLoggedModelTag: validate_can_delete_logged_model, + SetLoggedModelTags: validate_can_update_logged_model, + LogLoggedModelParamsRequest: validate_can_update_logged_model, +} + + +def get_logged_model_before_request_handler(request_class): + return LOGGED_MODEL_BEFORE_REQUEST_HANDLERS.get(request_class) + + +def _re_compile_path(path: str) -> re.Pattern: + """ + Convert a path with angle brackets to a regex pattern. For example, + "/api/2.0/experiments/" becomes "/api/2.0/experiments/([^/]+)". + """ + return re.compile(re.sub(r"<([^>]+)>", r"([^/]+)", path)) + + +LOGGED_MODEL_BEFORE_REQUEST_VALIDATORS = { + # Paths for logged models contains path parameters (e.g. /mlflow/logged-models/) + (_re_compile_path(http_path), method): handler + for http_path, handler, methods in get_endpoints(get_logged_model_before_request_handler) + for method in methods +} + +# Workspace RPC handlers (per decision WSAUTH-A: regex pattern matching like logged models) +WORKSPACE_BEFORE_REQUEST_HANDLERS = { + CreateWorkspace: validate_can_create_workspace, + GetWorkspace: validate_can_read_workspace, + ListWorkspaces: validate_can_list_workspaces, + UpdateWorkspace: validate_can_update_workspace, + DeleteWorkspace: validate_can_delete_workspace, +} + + +def get_workspace_before_request_handler(request_class): + return WORKSPACE_BEFORE_REQUEST_HANDLERS.get(request_class) + + +WORKSPACE_BEFORE_REQUEST_VALIDATORS = { + (_re_compile_path(http_path), method): handler + for http_path, handler, methods in get_endpoints(get_workspace_before_request_handler) + for method in methods + if handler is not None +} + + +# --------------------------------------------------------------------------- +# Workspace creation gating (per WSAUTH-F / WSAUTH-03) +# --------------------------------------------------------------------------- + +_WORKSPACE_GATED_CREATION_PATHS: set[tuple[str, str]] | None = None + + +def _get_workspace_gated_creation_paths() -> set[tuple[str, str]]: + """Lazily build the set of (path, method) pairs for workspace-gated creation.""" + global _WORKSPACE_GATED_CREATION_PATHS + if _WORKSPACE_GATED_CREATION_PATHS is None: + from mlflow.protos.service_pb2 import CreateExperiment + from mlflow.protos.model_registry_pb2 import CreateRegisteredModel + + paths = set() + for http_path, handler, methods in get_endpoints(lambda rc: rc if rc in (CreateExperiment, CreateRegisteredModel) else None): + if handler in (CreateExperiment, CreateRegisteredModel): + for method in methods: + paths.add((http_path, method)) + _WORKSPACE_GATED_CREATION_PATHS = paths + return _WORKSPACE_GATED_CREATION_PATHS + + +def _is_workspace_gated_creation(path: str, method: str) -> bool: + """Check if a request path/method corresponds to a workspace-gated creation endpoint.""" + return (path, method) in _get_workspace_gated_creation_paths() + + +def _get_proxy_artifact_validator(method: str, view_args: Optional[Dict[str, Any]]) -> Optional[Callable[[str], bool]]: + if view_args is None: + return validate_can_read_experiment_artifact_proxy # List + + return { + "GET": validate_can_read_experiment_artifact_proxy, # Download + "PUT": validate_can_update_experiment_artifact_proxy, # Upload + "DELETE": validate_can_delete_experiment_artifact_proxy, # Delete + }.get(method) + + +def _is_proxy_artifact_path(path: str) -> bool: + return path.startswith(f"{_REST_API_PATH_PREFIX}/mlflow-artifacts/artifacts/") + + +def _find_validator(req: Request) -> Optional[Callable[[str], bool]]: + """ + Finds the validator matching the request path and method. + """ + if "/mlflow/workspaces" in req.path: + # Workspace routes use path parameters (e.g. /mlflow/workspaces/) + validator = next( + (v for (pat, method), v in WORKSPACE_BEFORE_REQUEST_VALIDATORS.items() if pat.fullmatch(req.path) and method == req.method), + None, + ) + # Stash workspace name for after-request cascade delete (like gateway pattern) + if validator is not None and req.method == "DELETE": + from mlflow_oidc_auth.validators.workspace import ( + _extract_workspace_name_from_path, + ) + + ws_name = _extract_workspace_name_from_path() + if ws_name: + g._deleting_workspace_name = ws_name + return validator + if "/mlflow/logged-models" in req.path: + # logged model routes are not registered in the app + # so we need to check them manually + return next( + (v for (pat, method), v in LOGGED_MODEL_BEFORE_REQUEST_VALIDATORS.items() if pat.fullmatch(req.path) and method == req.method), + None, + ) + else: + return BEFORE_REQUEST_VALIDATORS.get((req.path, req.method)) + + +def before_request_hook(): + """Called before each request. If it did not return a response, + the view function for the matched route is called and returns a response""" + + if _is_unprotected_route(request.path): + return + + username, is_admin = _get_auth_context() + if username is None: + return responses.make_auth_required_response() + + logger.debug(f"Before request hook called for path: {request.path}, method: {request.method}, username: {username}, is admin: {is_admin}") + validator = _find_validator(request) + _stash_gateway_context(validator) + if is_admin: + return + # Workspace creation gating (per WSAUTH-F / WSAUTH-03) + if config.MLFLOW_ENABLE_WORKSPACES and _is_workspace_gated_creation(request.path, request.method): + from mlflow_oidc_auth.bridge.user import get_request_workspace + from mlflow_oidc_auth.utils.workspace_cache import ( + get_workspace_permission_cached, + ) + + workspace = get_request_workspace() + if workspace: + ws_perm = get_workspace_permission_cached(username, workspace) + # PATCH: can_manage → can_update + # OIDC_WORKSPACE_DEFAULT_PERMISSION auto-grants EDIT (can_update=True). + # Upstream erroneously required MANAGE (can_manage=True) for CreateExperiment/ + # CreateRegisteredModel, blocking all EDIT users. EDIT is the correct threshold. + if ws_perm is None or not ws_perm.can_update: + return responses.make_forbidden_response() + # authorization + if validator: + if not validator(username): + return responses.make_forbidden_response() + elif _is_proxy_artifact_path(request.path): + if validator := _get_proxy_artifact_validator(request.method, request.view_args): + if not validator(username): + return responses.make_forbidden_response() + + +before_request_hook = catch_mlflow_exception(before_request_hook) + + +def _stash_gateway_context(validator) -> None: + """Resolve and stash gateway resource names for after-request handlers. + + This must run for ALL users (including admins) because after-request + handlers need the old resource name to propagate permission changes + (renames) or clean up permission records (deletes). The before-request + validators run only for non-admin users and therefore cannot be relied + upon for stashing. + + The tracking store still has the old name/state at before-request time, + so ID-based resolution works correctly here. + """ + if validator is None: + return + + from mlflow_oidc_auth.validators.gateway import ( + _resolve_endpoint_name_from_id, + _resolve_secret_name_from_id, + _resolve_model_definition_name_from_id, + ) + + # --- Gateway endpoint: update (rename) or delete --- + if validator in ( + validate_can_update_gateway_endpoint, + validate_can_delete_gateway_endpoint, + ): + data = request.get_json(force=True, silent=True) or {} + endpoint_id = data.get("endpoint_id") + if endpoint_id: + name = _resolve_endpoint_name_from_id(endpoint_id) + if name: + if validator is validate_can_update_gateway_endpoint: + g._updating_gateway_endpoint_old_name = name + else: + g._deleting_gateway_endpoint_name = name + return + + # --- Gateway secret: delete --- + if validator is validate_can_delete_gateway_secret: + data = request.get_json(force=True, silent=True) or {} + secret_name = data.get("secret_name") + if not secret_name: + secret_id = data.get("secret_id") + if secret_id: + secret_name = _resolve_secret_name_from_id(secret_id) + if secret_name: + g._deleting_gateway_secret_name = secret_name + return + + # --- Gateway model definition: delete --- + if validator is validate_can_delete_gateway_model_definition: + data = request.get_json(force=True, silent=True) or {} + name = data.get("name") + if not name: + model_definition_id = data.get("model_definition_id") + if model_definition_id: + name = _resolve_model_definition_name_from_id(model_definition_id) + if name: + g._deleting_gateway_model_definition_name = name + return