diff --git a/sdk/python/feast/audit/__init__.py b/sdk/python/feast/audit/__init__.py new file mode 100644 index 00000000000..7b5d8ba7bb5 --- /dev/null +++ b/sdk/python/feast/audit/__init__.py @@ -0,0 +1,13 @@ +# Copyright 2025 The Feast Authors +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. diff --git a/sdk/python/feast/audit/audit_logger.py b/sdk/python/feast/audit/audit_logger.py new file mode 100644 index 00000000000..877380fe194 --- /dev/null +++ b/sdk/python/feast/audit/audit_logger.py @@ -0,0 +1,381 @@ +# Copyright 2025 The Feast Authors +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +""" +Structured audit logging for the Feast feature server. + +Emits JSONL audit events for MCP tool calls, REST requests, and +authentication/authorization decisions. Sensitive payloads (tokens, +entity rows, feature values) are never included. +""" + +import abc +import contextvars +import logging +import sys +import threading +import uuid +from datetime import datetime, timezone +from typing import TYPE_CHECKING, Any, Dict, List, Optional + +if TYPE_CHECKING: + from feast.infra.feature_servers.base_config import AuditLoggingConfig + +from pydantic import BaseModel, Field + +logger = logging.getLogger(__name__) + +# ContextVar used by the MCP audit wrapper to propagate a single request_id +# into the internal REST call so that ``mcp.tools.call`` and +# ``http.request`` events share the same identifier in SIEM. +mcp_audit_request_id: contextvars.ContextVar[Optional[str]] = contextvars.ContextVar( + "mcp_audit_request_id", default=None +) + + +# --------------------------------------------------------------------------- +# Audit event schema +# --------------------------------------------------------------------------- + + +class AuditPrincipal(BaseModel): + username: str = "" + roles: List[str] = Field(default_factory=list) + auth_type: str = "" + + +class AuditSource(BaseModel): + ip: str = "" + transport: str = "" + + +class AuditAction(BaseModel): + mcp_tool: str = "" + path: str = "" + + +class AuditResource(BaseModel): + type: str = "" + name: str = "" + actions: List[str] = Field(default_factory=list) + + +class AuditEvent(BaseModel): + """A single structured audit log entry. + + Fields follow the schema proposed in feast-dev/feast#6452. + No sensitive payloads (tokens, entity rows, feature values) are stored. + When OpenTelemetry is active, ``trace_id`` and ``span_id`` are populated + automatically for correlation with distributed traces. + """ + + event_type: str + timestamp: str = "" + request_id: str = "" + trace_id: Optional[str] = None + span_id: Optional[str] = None + jsonrpc_id: Optional[str] = None + principal: AuditPrincipal = Field(default_factory=AuditPrincipal) + source: AuditSource = Field(default_factory=AuditSource) + action: AuditAction = Field(default_factory=AuditAction) + resource: AuditResource = Field(default_factory=AuditResource) + outcome: str = "" + duration_ms: Optional[float] = None + detail: str = "" + + def to_jsonl(self) -> str: + return self.model_dump_json(exclude_none=True) + + +# --------------------------------------------------------------------------- +# Sink abstraction +# --------------------------------------------------------------------------- + + +class AuditSink(abc.ABC): + """Base class for audit event sinks.""" + + @abc.abstractmethod + def emit(self, event: AuditEvent) -> None: ... + + def close(self) -> None: + pass + + +class StdoutAuditSink(AuditSink): + """Write JSONL events to stdout (the default).""" + + def emit(self, event: AuditEvent) -> None: + sys.stdout.write(event.to_jsonl() + "\n") + sys.stdout.flush() + + +class FileAuditSink(AuditSink): + """Append JSONL events to a local file.""" + + def __init__(self, file_path: str) -> None: + self._file_path = file_path + self._fh = open(file_path, "a", buffering=1) + + def emit(self, event: AuditEvent) -> None: + self._fh.write(event.to_jsonl() + "\n") + + def close(self) -> None: + self._fh.close() + + +class LoggerAuditSink(AuditSink): + """Emit audit events through Python's ``logging`` module at INFO level.""" + + def __init__(self, logger_name: str = "feast.audit") -> None: + self._logger = logging.getLogger(logger_name) + + def emit(self, event: AuditEvent) -> None: + self._logger.info(event.to_jsonl()) + + +# --------------------------------------------------------------------------- +# AuditLogger — the main entry point +# --------------------------------------------------------------------------- + + +class AuditLogger: + """Central audit logger that routes events to the configured sink. + + Instantiate once during feature-server startup and share via + ``app.state.audit_logger``. + """ + + def __init__( + self, + sink: AuditSink, + *, + log_successful_reads: bool = True, + ) -> None: + self._sink = sink + self._log_successful_reads = log_successful_reads + self._lock = threading.Lock() + + # -- helpers ----------------------------------------------------------- + + @staticmethod + def new_request_id() -> str: + return str(uuid.uuid4()) + + @staticmethod + def _utcnow_iso() -> str: + return datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "Z" + + @staticmethod + def _inject_otel_context(event: AuditEvent) -> None: + """Populate trace_id/span_id from the active OpenTelemetry span, if any.""" + if event.trace_id is not None: + return + try: + from opentelemetry import trace as otel_trace + + span = otel_trace.get_current_span() + ctx = span.get_span_context() + if ctx and ctx.is_valid: + event.trace_id = format(ctx.trace_id, "032x") + event.span_id = format(ctx.span_id, "016x") + except ImportError: + pass + except Exception: + pass + + # MCP tool names and REST paths that correspond to read operations. + _READ_TOOL_NAMES = frozenset( + { + "get_online_features", + "retrieve_online_documents", + "get_historical_features", + } + ) + _READ_ACTIONS = frozenset({"READ_ONLINE", "READ_OFFLINE"}) + + @classmethod + def _is_read_event(cls, event: AuditEvent) -> bool: + """Return ``True`` when *event* represents a successful read. + + Checks both ``resource.actions`` (populated by REST middleware) and + ``action.mcp_tool`` (populated by the MCP handler wrapper) so that + suppression via ``log_successful_reads=False`` works for both layers. + """ + if event.event_type not in {"mcp.tools.call", "http.request"}: + return False + resource_actions = set(event.resource.actions) + if resource_actions and resource_actions.issubset(cls._READ_ACTIONS): + return True + if event.action.mcp_tool in cls._READ_TOOL_NAMES: + return True + return False + + # -- public API -------------------------------------------------------- + + def log(self, event: AuditEvent) -> None: + if not event.timestamp: + event.timestamp = self._utcnow_iso() + if not event.request_id: + event.request_id = self.new_request_id() + + self._inject_otel_context(event) + + if event.outcome == "success" and not self._log_successful_reads: + if self._is_read_event(event): + return + + try: + with self._lock: + self._sink.emit(event) + except Exception: + logger.exception("Failed to emit audit event") + + def log_mcp_call( + self, + *, + request_id: str, + tool_name: str, + path: str = "", + principal: Optional[AuditPrincipal] = None, + source: Optional[AuditSource] = None, + resource: Optional[AuditResource] = None, + outcome: str = "success", + duration_ms: Optional[float] = None, + detail: str = "", + ) -> None: + self.log( + AuditEvent( + event_type="mcp.tools.call", + request_id=request_id, + principal=principal or AuditPrincipal(), + source=source or AuditSource(transport="mcp-http"), + action=AuditAction(mcp_tool=tool_name, path=path), + resource=resource or AuditResource(), + outcome=outcome, + duration_ms=duration_ms, + detail=detail, + ) + ) + + def log_http_request( + self, + *, + request_id: str, + method: str, + path: str, + principal: Optional[AuditPrincipal] = None, + source: Optional[AuditSource] = None, + resource: Optional[AuditResource] = None, + outcome: str = "success", + duration_ms: Optional[float] = None, + status_code: int = 200, + ) -> None: + self.log( + AuditEvent( + event_type="http.request", + request_id=request_id, + principal=principal or AuditPrincipal(), + source=source or AuditSource(transport="http"), + action=AuditAction(path=path), + resource=resource or AuditResource(), + outcome=outcome, + duration_ms=duration_ms, + detail=f"{method} {path} -> {status_code}", + ) + ) + + def log_authn( + self, + *, + request_id: str, + outcome: str, + principal: Optional[AuditPrincipal] = None, + source: Optional[AuditSource] = None, + detail: str = "", + ) -> None: + event_type = "authn.success" if outcome == "success" else "authn.failure" + self.log( + AuditEvent( + event_type=event_type, + request_id=request_id, + principal=principal or AuditPrincipal(), + source=source or AuditSource(), + outcome=outcome, + detail=detail, + ) + ) + + def log_authz( + self, + *, + request_id: str, + outcome: str, + principal: Optional[AuditPrincipal] = None, + resource: Optional[AuditResource] = None, + detail: str = "", + ) -> None: + self.log( + AuditEvent( + event_type="authz.decision", + request_id=request_id, + principal=principal or AuditPrincipal(), + resource=resource or AuditResource(), + outcome=outcome, + detail=detail, + ) + ) + + def close(self) -> None: + self._sink.close() + + +# --------------------------------------------------------------------------- +# Factory +# --------------------------------------------------------------------------- + +_SINK_FACTORIES = { + "stdout": lambda cfg: StdoutAuditSink(), + "file": lambda cfg: FileAuditSink(cfg.get("file_path", "feast_audit.log")), + "logger": lambda cfg: LoggerAuditSink(cfg.get("logger_name", "feast.audit")), +} + + +def create_audit_logger_from_config( + audit_cfg: Optional["AuditLoggingConfig"], +) -> Optional[AuditLogger]: + """Build an ``AuditLogger`` from an ``AuditLoggingConfig`` pydantic model. + + Returns ``None`` when audit logging is disabled. + """ + if audit_cfg is None or not getattr(audit_cfg, "enabled", False): + return None + + sink_type = getattr(audit_cfg, "sink", "stdout") + raw: Dict[str, Any] = {} + if hasattr(audit_cfg, "file_path"): + raw["file_path"] = audit_cfg.file_path + if hasattr(audit_cfg, "logger_name"): + raw["logger_name"] = audit_cfg.logger_name + + factory = _SINK_FACTORIES.get(sink_type) + if factory is None: + logger.warning("Unknown audit sink %r, falling back to stdout", sink_type) + factory = _SINK_FACTORIES["stdout"] + + sink = factory(raw) + return AuditLogger( + sink, + log_successful_reads=getattr(audit_cfg, "log_successful_reads", True), + ) diff --git a/sdk/python/feast/audit/audit_middleware.py b/sdk/python/feast/audit/audit_middleware.py new file mode 100644 index 00000000000..d56ca01d553 --- /dev/null +++ b/sdk/python/feast/audit/audit_middleware.py @@ -0,0 +1,144 @@ +# Copyright 2025 The Feast Authors +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +""" +FastAPI middleware for structured audit logging of REST endpoints. + +``AuditLoggingMiddleware`` logs every HTTP request/response on the REST +endpoints (``/get-online-features``, ``/push``, etc.) with principal, +resource, outcome, and duration. + +MCP tool-call auditing is handled at the protocol layer by wrapping +the ``tools/call`` handler inside ``add_mcp_support_to_app()`` (see +``feast.infra.mcp_servers.mcp_server``). + +The middleware is added only when ``audit_logging.enabled`` is ``true`` +in the feature-server configuration. +""" + +import logging +import time +from typing import Optional + +from starlette.middleware.base import BaseHTTPMiddleware +from starlette.requests import Request +from starlette.responses import Response + +from feast.audit.audit_logger import ( + AuditLogger, + AuditPrincipal, + AuditResource, + AuditSource, + mcp_audit_request_id, +) + +logger = logging.getLogger(__name__) + +# REST paths that correspond to read/write actions +_PATH_RESOURCE_MAP = { + "/get-online-features": ("feature_service", ["READ_ONLINE"]), + "/retrieve-online-documents": ("feature_service", ["READ_ONLINE"]), + "/push": ("push_source", ["WRITE_ONLINE", "WRITE_OFFLINE"]), + "/write-to-online-store": ("feature_view", ["WRITE_ONLINE"]), + "/materialize": ("feature_view", ["WRITE_ONLINE"]), + "/materialize-incremental": ("feature_view", ["WRITE_ONLINE"]), +} + + +def _extract_client_ip(request: Request) -> str: + forwarded = request.headers.get("x-forwarded-for", "") + if forwarded: + return forwarded.split(",")[0].strip() + if request.client: + return request.client.host + return "" + + +def _principal_from_request(request: Request) -> AuditPrincipal: + """Build an ``AuditPrincipal`` from the security manager's current user.""" + try: + from feast.permissions.security_manager import get_security_manager + + sm = get_security_manager() + if sm and sm.current_user: + user = sm.current_user + return AuditPrincipal( + username=user.username, + roles=list(user.roles) if user.roles else [], + auth_type=request.headers.get("x-feast-auth-type", ""), + ) + except Exception: + pass + return AuditPrincipal() + + +class AuditLoggingMiddleware(BaseHTTPMiddleware): + """Emit ``http.request`` audit events for REST endpoints.""" + + async def dispatch(self, request: Request, call_next): # type: ignore[override] + audit: Optional[AuditLogger] = getattr(request.app.state, "audit_logger", None) + if audit is None: + return await call_next(request) + + path = request.url.path + # Skip health and static endpoints + if path in ("/health", "/docs", "/openapi.json") or path.startswith("/static"): + return await call_next(request) + + # Skip MCP transport endpoints — tool-call auditing is handled at the + # protocol layer in mcp_server._wrap_call_tool_handler(). + if path.startswith("/mcp"): + return await call_next(request) + + # If this request was triggered by an MCP tool call, reuse the same + # request_id so both events correlate in SIEM. + propagated_id = mcp_audit_request_id.get() + request_id = ( + propagated_id + or request.headers.get("x-request-id") + or audit.new_request_id() + ) + request.state.audit_request_id = request_id + + start = time.monotonic() + response: Response + outcome = "success" + status_code = 200 + try: + response = await call_next(request) + status_code = response.status_code + if status_code >= 400: + outcome = "failure" + except Exception: + outcome = "error" + status_code = 500 + raise + finally: + duration_ms = (time.monotonic() - start) * 1000.0 + resource_info = _PATH_RESOURCE_MAP.get(path, ("", [])) + audit.log_http_request( + request_id=request_id, + method=request.method, + path=path, + principal=_principal_from_request(request), + source=AuditSource(ip=_extract_client_ip(request), transport="http"), + resource=AuditResource( + type=resource_info[0], actions=list(resource_info[1]) + ), + outcome=outcome, + duration_ms=round(duration_ms, 2), + status_code=status_code, + ) + + return response diff --git a/sdk/python/feast/feature_server.py b/sdk/python/feast/feature_server.py index e42e28d6db9..eb85234822f 100644 --- a/sdk/python/feast/feature_server.py +++ b/sdk/python/feast/feature_server.py @@ -587,11 +587,27 @@ def async_refresh(): active_timer = threading.Timer(registry_ttl_sec, async_refresh) active_timer.start() + # --- Audit logging setup --- + audit_logging_cfg = getattr(fs_cfg, "audit_logging", None) + audit_logger_instance = None + if audit_logging_cfg is not None and getattr(audit_logging_cfg, "enabled", False): + from feast.audit.audit_logger import create_audit_logger_from_config + + audit_logger_instance = create_audit_logger_from_config(audit_logging_cfg) + if audit_logger_instance: + logger.info( + "Structured audit logging is ENABLED (sink=%s)", + getattr(audit_logging_cfg, "sink", "stdout"), + ) + @asynccontextmanager async def lifespan(app: FastAPI): # Load static artifacts before initializing store await load_static_artifacts(app, store) + if audit_logger_instance is not None: + app.state.audit_logger = audit_logger_instance + await store.initialize() async_refresh() try: @@ -603,10 +619,19 @@ async def lifespan(app: FastAPI): # wait=False: do not block process exit on in-flight materialize # (same fire-and-forget contract as returning 202 mid-job). materialize_executor.shutdown(wait=False) + if audit_logger_instance is not None: + audit_logger_instance.close() await store.close() app = FastAPI(lifespan=lifespan) + # Add audit logging middleware when enabled (REST only; + # MCP audit is handled at the protocol layer in mcp_server.py) + if audit_logger_instance is not None: + from feast.audit.audit_middleware import AuditLoggingMiddleware + + app.add_middleware(AuditLoggingMiddleware) + @app.post( "/get-online-features", dependencies=[Depends(inject_user_details)], @@ -1153,12 +1178,12 @@ async def websocket_endpoint(websocket: WebSocket): app.mount("/static", StaticFiles(directory=static_dir), name="static") # Add MCP support if enabled in feature server configuration - _add_mcp_support_if_enabled(app, store) + _add_mcp_support_if_enabled(app, store, audit_logger_instance) return app -def _add_mcp_support_if_enabled(app, store: "feast.FeatureStore"): +def _add_mcp_support_if_enabled(app, store: "feast.FeatureStore", audit_logger=None): """Add MCP support to the FastAPI app if enabled in configuration.""" mcp_transport_not_supported_error = None try: @@ -1180,7 +1205,9 @@ def _add_mcp_support_if_enabled(app, store: "feast.FeatureStore"): logger.error(f"Error checking/adding MCP support: {e}") return - mcp_server = add_mcp_support_to_app(app, store, store.config.feature_server) + mcp_server = add_mcp_support_to_app( + app, store, store.config.feature_server, audit_logger=audit_logger + ) if mcp_server: logger.info("MCP support has been enabled for the Feast feature server") diff --git a/sdk/python/feast/infra/feature_servers/base_config.py b/sdk/python/feast/infra/feature_servers/base_config.py index 14ad2fe505e..4a534210712 100644 --- a/sdk/python/feast/infra/feature_servers/base_config.py +++ b/sdk/python/feast/infra/feature_servers/base_config.py @@ -11,7 +11,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -from typing import Optional +from typing import Literal, Optional from pydantic import StrictBool, StrictInt @@ -94,6 +94,27 @@ class MetricsConfig(FeastConfigBaseModel): identity, entity keys, feature views, row counts, and latency.""" +class AuditLoggingConfig(FeastConfigBaseModel): + """Structured audit logging configuration for the feature server. + + Emits JSONL audit events for MCP tool calls, REST requests, + and authentication/authorization decisions. + """ + + enabled: StrictBool = False + """Whether structured audit logging is enabled.""" + + sink: Literal["stdout", "file", "logger"] = "stdout" + """Audit event sink: ``stdout``, ``file``, or ``logger``.""" + + file_path: str = "feast_audit.log" + """File path when ``sink`` is ``file``.""" + + log_successful_reads: StrictBool = True + """Emit audit events for successful read operations. Set to ``False`` + to reduce log volume in high-throughput read-heavy deployments.""" + + class BaseFeatureServerConfig(FeastConfigBaseModel): """Base Feature Server config that should be extended""" @@ -107,6 +128,10 @@ class BaseFeatureServerConfig(FeastConfigBaseModel): feature_logging: Optional[FeatureLoggingConfig] = None """ Feature logging configuration """ + audit_logging: Optional[AuditLoggingConfig] = None + """Structured audit logging configuration. Emits JSONL audit events + for MCP tool calls, REST requests, and auth decisions.""" + offline_push_batching_enabled: Optional[StrictBool] = None """Whether to batch writes to the offline store via the `/push` endpoint.""" diff --git a/sdk/python/feast/infra/mcp_servers/mcp_server.py b/sdk/python/feast/infra/mcp_servers/mcp_server.py index 2be788148cd..dc1cbd156e1 100644 --- a/sdk/python/feast/infra/mcp_servers/mcp_server.py +++ b/sdk/python/feast/infra/mcp_servers/mcp_server.py @@ -3,9 +3,15 @@ This module provides MCP support for Feast by integrating with fastapi_mcp to expose Feast functionality through the Model Context Protocol. + +When audit logging is enabled, the ``CallToolRequest`` handler on the +low-level MCP ``Server`` is wrapped so that every tool invocation is +logged with typed tool name, outcome, duration, and principal — without +parsing raw JSON-RPC bodies. """ import logging +import time from typing import TYPE_CHECKING, Any, Dict, Optional, Set if TYPE_CHECKING: @@ -23,9 +29,13 @@ "Install it with: pip install fastapi_mcp" ) MCP_AVAILABLE = False - # Create placeholder classes for testing FastApiMCP = None +try: + from mcp.types import CallToolRequest as _CallToolRequest +except ImportError: + _CallToolRequest = None # type: ignore[assignment,misc] + class McpTransportNotSupportedError(RuntimeError): pass @@ -99,7 +109,10 @@ def _patch_fastapi_mcp_schema_resolver() -> None: def add_mcp_support_to_app( - app, store: "FeatureStore", config + app, + store: "FeatureStore", + config, + audit_logger: Optional[Any] = None, ) -> Optional["FastApiMCP"]: """Add MCP support to the FastAPI app if enabled in configuration.""" if not MCP_AVAILABLE: @@ -135,11 +148,13 @@ def add_mcp_support_to_app( ) mcp.mount() else: - # Defensive guard for programmatic callers. raise McpTransportNotSupportedError( f"Unsupported mcp_transport={transport!r}. Expected 'sse' or 'http'." ) + if audit_logger is not None: + _wrap_call_tool_handler(mcp, audit_logger) + logger.info( "MCP support has been enabled for the Feast feature server at /mcp endpoint" ) @@ -155,3 +170,121 @@ def add_mcp_support_to_app( except Exception as e: logger.error(f"Failed to initialize MCP integration: {e}", exc_info=True) return None + + +# --------------------------------------------------------------------------- +# Audit-logging wrapper for the MCP tools/call handler +# --------------------------------------------------------------------------- + + +def _get_call_tool_handler_key() -> Any: + """Return the dict key used for CallToolRequest in ``server.request_handlers``. + + mcp 1.x uses the ``CallToolRequest`` *class* as the key in + ``server.request_handlers``. + """ + if _CallToolRequest is not None: + return _CallToolRequest + return None + + +def _principal_from_mcp_context(server: Any) -> Any: + """Extract an ``AuditPrincipal`` from the MCP server's request context. + + In mcp 1.x the request context is a ``ContextVar`` accessed via + ``server.request_context``. The ``.request`` attribute carries the + original Starlette/FastAPI ``Request`` that ``fastapi_mcp`` injects + through ``ServerMessageMetadata(request_context=request)``. + """ + from feast.audit.audit_logger import AuditPrincipal + + try: + ctx = server.request_context + request = getattr(ctx, "request", None) + if request is None: + return AuditPrincipal() + headers: dict[str, str] = {} + if hasattr(request, "headers"): + headers = dict(request.headers) + auth_type = headers.get("x-feast-auth-type", "") + has_auth = bool(headers.get("authorization", "")) + return AuditPrincipal( + username="(authenticated)" if has_auth else "", + auth_type=auth_type, + ) + except Exception: + return AuditPrincipal() + + +def _wrap_call_tool_handler(mcp: "FastApiMCP", audit: Any) -> None: + """Wrap the MCP server's ``CallToolRequest`` handler with audit logging. + + In mcp 1.x the handler lives at + ``server.request_handlers[CallToolRequest]`` and has the signature + ``async def handler(req: CallToolRequest) -> ServerResult``. The + JSON-RPC request_id is available on ``server.request_context``. + """ + from feast.audit.audit_logger import AuditAction, AuditEvent, AuditSource + + handler_key = _get_call_tool_handler_key() + handlers = getattr(mcp.server, "request_handlers", None) + if handlers is None: + logger.warning("Cannot wrap MCP call_tool handler: request_handlers not found") + return + + if handler_key is None or handler_key not in handlers: + logger.debug("No CallToolRequest handler registered; skipping audit wrapper") + return + + original = handlers[handler_key] + + async def audited_call_tool(req: Any) -> Any: + from feast.audit.audit_logger import mcp_audit_request_id + + params = getattr(req, "params", None) + tool_name = getattr(params, "name", "") if params else "" + request_id = audit.new_request_id() + + jsonrpc_id: Optional[str] = None + try: + ctx = mcp.server.request_context + if hasattr(ctx, "request_id"): + jsonrpc_id = str(ctx.request_id) + except LookupError: + pass + + token = mcp_audit_request_id.set(request_id) + start = time.monotonic() + outcome = "success" + error_detail = "" + try: + result = await original(req) + # The MCP SDK returns ServerResult(root=CallToolResult(...)); + # isError lives on the inner CallToolResult, not the wrapper. + inner = getattr(result, "root", result) + if getattr(inner, "isError", False): + outcome = "mcp_error" + return result + except Exception as exc: + outcome = "error" + error_detail = str(exc)[:200] + raise + finally: + duration_ms = (time.monotonic() - start) * 1000.0 + mcp_audit_request_id.reset(token) + principal = _principal_from_mcp_context(mcp.server) + audit.log( + AuditEvent( + event_type="mcp.tools.call", + request_id=request_id, + jsonrpc_id=jsonrpc_id, + principal=principal, + source=AuditSource(transport="mcp-http"), + action=AuditAction(mcp_tool=tool_name), + outcome=outcome, + duration_ms=round(duration_ms, 2), + detail=error_detail, + ) + ) + + handlers[handler_key] = audited_call_tool diff --git a/sdk/python/tests/unit/audit/__init__.py b/sdk/python/tests/unit/audit/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/sdk/python/tests/unit/audit/test_audit_logger.py b/sdk/python/tests/unit/audit/test_audit_logger.py new file mode 100644 index 00000000000..34b0d57bd74 --- /dev/null +++ b/sdk/python/tests/unit/audit/test_audit_logger.py @@ -0,0 +1,434 @@ +import json +import os +import tempfile +import threading +import unittest +from unittest.mock import MagicMock, patch + +from feast.audit.audit_logger import ( + AuditEvent, + AuditLogger, + AuditPrincipal, + AuditResource, + AuditSink, + AuditSource, + FileAuditSink, + LoggerAuditSink, + StdoutAuditSink, + create_audit_logger_from_config, +) + + +class InMemorySink(AuditSink): + """Test sink that captures events in a list.""" + + def __init__(self): + self.events: list[AuditEvent] = [] + + def emit(self, event: AuditEvent) -> None: + self.events.append(event) + + +class TestAuditEvent(unittest.TestCase): + def test_event_to_jsonl(self): + event = AuditEvent( + event_type="mcp.tools.call", + timestamp="2026-05-28T12:00:00.000Z", + request_id="abc-123", + principal=AuditPrincipal( + username="jane@co.com", roles=["reader"], auth_type="oidc" + ), + source=AuditSource(ip="10.0.0.1", transport="mcp-http"), + outcome="success", + duration_ms=42.0, + ) + line = event.to_jsonl() + parsed = json.loads(line) + + self.assertEqual(parsed["event_type"], "mcp.tools.call") + self.assertEqual(parsed["principal"]["username"], "jane@co.com") + self.assertEqual(parsed["source"]["ip"], "10.0.0.1") + self.assertEqual(parsed["duration_ms"], 42.0) + self.assertNotIn("\n", line) + + def test_event_excludes_none_duration(self): + event = AuditEvent(event_type="authn.success") + parsed = json.loads(event.to_jsonl()) + self.assertNotIn("duration_ms", parsed) + + def test_event_includes_trace_id_and_span_id(self): + event = AuditEvent( + event_type="mcp.tools.call", + trace_id="0" * 32, + span_id="f" * 16, + ) + parsed = json.loads(event.to_jsonl()) + self.assertEqual(parsed["trace_id"], "0" * 32) + self.assertEqual(parsed["span_id"], "f" * 16) + + def test_event_excludes_none_trace_fields(self): + event = AuditEvent(event_type="test") + parsed = json.loads(event.to_jsonl()) + self.assertNotIn("trace_id", parsed) + self.assertNotIn("span_id", parsed) + self.assertNotIn("jsonrpc_id", parsed) + + def test_event_includes_jsonrpc_id(self): + event = AuditEvent(event_type="mcp.tools.call", jsonrpc_id="42") + parsed = json.loads(event.to_jsonl()) + self.assertEqual(parsed["jsonrpc_id"], "42") + + +class TestAuditLogger(unittest.TestCase): + def test_log_populates_timestamp_and_request_id(self): + sink = InMemorySink() + al = AuditLogger(sink) + al.log(AuditEvent(event_type="test")) + + self.assertEqual(len(sink.events), 1) + event = sink.events[0] + self.assertTrue(event.timestamp.endswith("Z")) + self.assertTrue(len(event.request_id) > 0) + + def test_log_mcp_call(self): + sink = InMemorySink() + al = AuditLogger(sink) + al.log_mcp_call( + request_id="r1", + tool_name="get_online_features", + path="/get-online-features", + outcome="success", + duration_ms=10.5, + ) + + self.assertEqual(len(sink.events), 1) + event = sink.events[0] + self.assertEqual(event.event_type, "mcp.tools.call") + self.assertEqual(event.action.mcp_tool, "get_online_features") + self.assertEqual(event.outcome, "success") + + def test_log_http_request(self): + sink = InMemorySink() + al = AuditLogger(sink) + al.log_http_request( + request_id="r2", + method="POST", + path="/push", + status_code=200, + ) + + self.assertEqual(len(sink.events), 1) + event = sink.events[0] + self.assertEqual(event.event_type, "http.request") + self.assertIn("POST /push -> 200", event.detail) + + def test_log_authn_success(self): + sink = InMemorySink() + al = AuditLogger(sink) + al.log_authn(request_id="r3", outcome="success") + self.assertEqual(sink.events[0].event_type, "authn.success") + + def test_log_authn_failure(self): + sink = InMemorySink() + al = AuditLogger(sink) + al.log_authn(request_id="r4", outcome="failure", detail="bad token") + event = sink.events[0] + self.assertEqual(event.event_type, "authn.failure") + self.assertEqual(event.detail, "bad token") + + def test_log_authz_decision(self): + sink = InMemorySink() + al = AuditLogger(sink) + al.log_authz( + request_id="r5", + outcome="denied", + resource=AuditResource( + type="feature_service", name="driver_fs", actions=["READ_ONLINE"] + ), + ) + event = sink.events[0] + self.assertEqual(event.event_type, "authz.decision") + self.assertEqual(event.resource.name, "driver_fs") + + def test_log_successful_reads_suppressed(self): + sink = InMemorySink() + al = AuditLogger(sink, log_successful_reads=False) + al.log_http_request( + request_id="r6", + method="POST", + path="/get-online-features", + resource=AuditResource(type="feature_service", actions=["READ_ONLINE"]), + outcome="success", + status_code=200, + ) + # Successful read should be suppressed + self.assertEqual(len(sink.events), 0) + + def test_log_successful_reads_not_suppressed_for_writes(self): + sink = InMemorySink() + al = AuditLogger(sink, log_successful_reads=False) + al.log_http_request( + request_id="r7", + method="POST", + path="/push", + resource=AuditResource(type="push_source", actions=["WRITE_ONLINE"]), + outcome="success", + status_code=200, + ) + self.assertEqual(len(sink.events), 1) + + def test_log_successful_mcp_read_suppressed(self): + """MCP tool-call events for read tools are suppressed when + log_successful_reads=False, even without resource.actions.""" + sink = InMemorySink() + al = AuditLogger(sink, log_successful_reads=False) + al.log_mcp_call( + request_id="r-mcp", + tool_name="get_online_features", + outcome="success", + ) + self.assertEqual(len(sink.events), 0) + + def test_log_successful_mcp_write_not_suppressed(self): + """MCP tool-call events for non-read tools are always logged.""" + sink = InMemorySink() + al = AuditLogger(sink, log_successful_reads=False) + al.log_mcp_call( + request_id="r-mcp-w", + tool_name="push", + outcome="success", + ) + self.assertEqual(len(sink.events), 1) + + def test_log_failed_reads_not_suppressed(self): + sink = InMemorySink() + al = AuditLogger(sink, log_successful_reads=False) + al.log_http_request( + request_id="r8", + method="POST", + path="/get-online-features", + resource=AuditResource(type="feature_service", actions=["READ_ONLINE"]), + outcome="failure", + status_code=500, + ) + self.assertEqual(len(sink.events), 1) + + def test_emit_exception_does_not_raise(self): + bad_sink = MagicMock(spec=AuditSink) + bad_sink.emit.side_effect = RuntimeError("disk full") + al = AuditLogger(bad_sink) + al.log(AuditEvent(event_type="test")) + + def test_otel_context_injected_when_active(self): + import sys + import types + + mock_ctx = MagicMock() + mock_ctx.is_valid = True + mock_ctx.trace_id = 0xABCDEF1234567890ABCDEF1234567890 + mock_ctx.span_id = 0x1234567890ABCDEF + + mock_span = MagicMock() + mock_span.get_span_context.return_value = mock_ctx + + mock_otel_trace = types.ModuleType("opentelemetry.trace") + mock_otel_trace.get_current_span = lambda: mock_span # type: ignore[attr-defined] + + mock_otel = types.ModuleType("opentelemetry") + mock_otel.trace = mock_otel_trace # type: ignore[attr-defined] + + saved_otel = sys.modules.get("opentelemetry") + saved_trace = sys.modules.get("opentelemetry.trace") + sys.modules["opentelemetry"] = mock_otel + sys.modules["opentelemetry.trace"] = mock_otel_trace + try: + sink = InMemorySink() + al = AuditLogger(sink) + al.log(AuditEvent(event_type="test")) + event = sink.events[0] + self.assertEqual(event.trace_id, "abcdef1234567890abcdef1234567890") + self.assertEqual(event.span_id, "1234567890abcdef") + finally: + if saved_otel is None: + sys.modules.pop("opentelemetry", None) + else: + sys.modules["opentelemetry"] = saved_otel + if saved_trace is None: + sys.modules.pop("opentelemetry.trace", None) + else: + sys.modules["opentelemetry.trace"] = saved_trace + + def test_otel_not_available_no_crash(self): + sink = InMemorySink() + al = AuditLogger(sink) + al.log(AuditEvent(event_type="test")) + event = sink.events[0] + self.assertIsNone(event.trace_id) + self.assertIsNone(event.span_id) + + def test_otel_skipped_when_trace_id_already_set(self): + sink = InMemorySink() + al = AuditLogger(sink) + al.log(AuditEvent(event_type="test", trace_id="preexisting")) + event = sink.events[0] + self.assertEqual(event.trace_id, "preexisting") + + def test_audit_logger_has_lock(self): + sink = InMemorySink() + al = AuditLogger(sink) + self.assertIsInstance(al._lock, type(threading.Lock())) + + +class TestStdoutSink(unittest.TestCase): + @patch("sys.stdout") + def test_emit_writes_to_stdout(self, mock_stdout): + sink = StdoutAuditSink() + event = AuditEvent(event_type="test", timestamp="t", request_id="r") + sink.emit(event) + mock_stdout.write.assert_called_once() + written = mock_stdout.write.call_args[0][0] + self.assertIn('"event_type":"test"', written) + self.assertTrue(written.endswith("\n")) + + +class TestFileAuditSink(unittest.TestCase): + def test_emit_appends_to_file(self): + with tempfile.NamedTemporaryFile(mode="w", suffix=".log", delete=False) as tmp: + path = tmp.name + + try: + sink = FileAuditSink(path) + event = AuditEvent(event_type="test", timestamp="t", request_id="r") + sink.emit(event) + sink.close() + + with open(path) as f: + lines = f.readlines() + self.assertEqual(len(lines), 1) + parsed = json.loads(lines[0]) + self.assertEqual(parsed["event_type"], "test") + finally: + os.unlink(path) + + +class TestLoggerAuditSink(unittest.TestCase): + @patch("logging.getLogger") + def test_emit_logs_at_info(self, mock_get_logger): + mock_logger = MagicMock() + mock_get_logger.return_value = mock_logger + + sink = LoggerAuditSink("feast.audit") + event = AuditEvent(event_type="test", timestamp="t", request_id="r") + sink.emit(event) + mock_logger.info.assert_called_once() + + +class TestCreateAuditLoggerFromConfig(unittest.TestCase): + def test_returns_none_when_disabled(self): + from types import SimpleNamespace + + cfg = SimpleNamespace(enabled=False) + self.assertIsNone(create_audit_logger_from_config(cfg)) + + def test_returns_none_when_none(self): + self.assertIsNone(create_audit_logger_from_config(None)) + + def test_creates_stdout_logger(self): + from types import SimpleNamespace + + cfg = SimpleNamespace(enabled=True, sink="stdout", log_successful_reads=True) + al = create_audit_logger_from_config(cfg) + self.assertIsNotNone(al) + self.assertIsInstance(al._sink, StdoutAuditSink) + + def test_creates_file_logger(self): + from types import SimpleNamespace + + with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as tmp: + path = tmp.name + + try: + cfg = SimpleNamespace( + enabled=True, sink="file", file_path=path, log_successful_reads=True + ) + al = create_audit_logger_from_config(cfg) + self.assertIsNotNone(al) + self.assertIsInstance(al._sink, FileAuditSink) + al.close() + finally: + os.unlink(path) + + def test_creates_logger_sink(self): + from types import SimpleNamespace + + cfg = SimpleNamespace(enabled=True, sink="logger", log_successful_reads=True) + al = create_audit_logger_from_config(cfg) + self.assertIsNotNone(al) + self.assertIsInstance(al._sink, LoggerAuditSink) + + def test_unknown_sink_falls_back_to_stdout(self): + from types import SimpleNamespace + + cfg = SimpleNamespace(enabled=True, sink="kafka", log_successful_reads=True) + al = create_audit_logger_from_config(cfg) + self.assertIsNotNone(al) + self.assertIsInstance(al._sink, StdoutAuditSink) + + +class TestAuditLoggingConfig(unittest.TestCase): + def test_config_defaults(self): + from feast.infra.feature_servers.base_config import AuditLoggingConfig + + cfg = AuditLoggingConfig() + self.assertFalse(cfg.enabled) + self.assertEqual(cfg.sink, "stdout") + self.assertEqual(cfg.file_path, "feast_audit.log") + self.assertTrue(cfg.log_successful_reads) + + def test_config_custom(self): + from feast.infra.feature_servers.base_config import AuditLoggingConfig + + cfg = AuditLoggingConfig( + enabled=True, + sink="file", + file_path="/var/log/feast_audit.jsonl", + log_successful_reads=False, + ) + self.assertTrue(cfg.enabled) + self.assertEqual(cfg.sink, "file") + self.assertEqual(cfg.file_path, "/var/log/feast_audit.jsonl") + self.assertFalse(cfg.log_successful_reads) + + def test_base_feature_server_config_includes_audit_logging(self): + from feast.infra.feature_servers.base_config import ( + AuditLoggingConfig, + BaseFeatureServerConfig, + ) + + cfg = BaseFeatureServerConfig( + audit_logging=AuditLoggingConfig(enabled=True, sink="stdout") + ) + self.assertIsNotNone(cfg.audit_logging) + self.assertTrue(cfg.audit_logging.enabled) + + def test_invalid_sink_rejected(self): + from pydantic import ValidationError + + from feast.infra.feature_servers.base_config import AuditLoggingConfig + + with self.assertRaises(ValidationError): + AuditLoggingConfig(enabled=True, sink="kafka") + + def test_mcp_config_includes_audit_logging(self): + from feast.infra.feature_servers.base_config import AuditLoggingConfig + from feast.infra.mcp_servers.mcp_config import McpFeatureServerConfig + + cfg = McpFeatureServerConfig( + mcp_enabled=True, + audit_logging=AuditLoggingConfig( + enabled=True, sink="file", file_path="/tmp/audit.log" + ), + ) + self.assertIsNotNone(cfg.audit_logging) + self.assertTrue(cfg.audit_logging.enabled) + self.assertEqual(cfg.audit_logging.sink, "file") diff --git a/sdk/python/tests/unit/audit/test_audit_middleware.py b/sdk/python/tests/unit/audit/test_audit_middleware.py new file mode 100644 index 00000000000..dd9feca7f8b --- /dev/null +++ b/sdk/python/tests/unit/audit/test_audit_middleware.py @@ -0,0 +1,132 @@ +import unittest + +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from feast.audit.audit_logger import ( + AuditEvent, + AuditLogger, + AuditSink, +) +from feast.audit.audit_middleware import AuditLoggingMiddleware + + +class InMemorySink(AuditSink): + def __init__(self): + self.events: list[AuditEvent] = [] + + def emit(self, event: AuditEvent) -> None: + self.events.append(event) + + +def _make_app(audit_logger=None): + """Create a minimal FastAPI app with audit middleware for testing.""" + app = FastAPI() + app.state.audit_logger = audit_logger + + app.add_middleware(AuditLoggingMiddleware) + + @app.post("/get-online-features") + async def get_online_features(): + return {"result": "ok"} + + @app.post("/push") + async def push(): + return {"result": "ok"} + + @app.get("/health") + async def health(): + return {"status": "ok"} + + @app.post("/mcp") + async def mcp(): + return {"jsonrpc": "2.0", "result": "ok", "id": 1} + + @app.post("/error-endpoint") + async def error_endpoint(): + raise ValueError("test error") + + return app + + +class TestAuditLoggingMiddleware(unittest.TestCase): + def test_logs_http_request(self): + sink = InMemorySink() + audit = AuditLogger(sink) + app = _make_app(audit) + client = TestClient(app, raise_server_exceptions=False) + + resp = client.post("/get-online-features") + self.assertEqual(resp.status_code, 200) + + http_events = [e for e in sink.events if e.event_type == "http.request"] + self.assertEqual(len(http_events), 1) + event = http_events[0] + self.assertEqual(event.outcome, "success") + self.assertIn("/get-online-features", event.detail) + self.assertIsNotNone(event.duration_ms) + self.assertGreaterEqual(event.duration_ms, 0) + + def test_skips_health_endpoint(self): + sink = InMemorySink() + audit = AuditLogger(sink) + app = _make_app(audit) + client = TestClient(app, raise_server_exceptions=False) + + client.get("/health") + http_events = [e for e in sink.events if e.event_type == "http.request"] + self.assertEqual(len(http_events), 0) + + def test_skips_mcp_endpoint(self): + sink = InMemorySink() + audit = AuditLogger(sink) + app = _make_app(audit) + client = TestClient(app, raise_server_exceptions=False) + + client.post("/mcp", json={"jsonrpc": "2.0", "method": "tools/list"}) + http_events = [e for e in sink.events if e.event_type == "http.request"] + self.assertEqual(len(http_events), 0) + + def test_logs_failure_on_error(self): + sink = InMemorySink() + audit = AuditLogger(sink) + app = _make_app(audit) + client = TestClient(app, raise_server_exceptions=False) + + resp = client.post("/error-endpoint") + self.assertEqual(resp.status_code, 500) + + http_events = [e for e in sink.events if e.event_type == "http.request"] + self.assertEqual(len(http_events), 1) + self.assertEqual(http_events[0].outcome, "error") + + def test_no_logging_when_audit_logger_is_none(self): + app = _make_app(audit_logger=None) + client = TestClient(app, raise_server_exceptions=False) + + resp = client.post("/get-online-features") + self.assertEqual(resp.status_code, 200) + + def test_uses_x_request_id_header(self): + sink = InMemorySink() + audit = AuditLogger(sink) + app = _make_app(audit) + client = TestClient(app, raise_server_exceptions=False) + + client.post( + "/get-online-features", + headers={"x-request-id": "custom-id-123"}, + ) + http_events = [e for e in sink.events if e.event_type == "http.request"] + self.assertEqual(http_events[0].request_id, "custom-id-123") + + def test_resource_mapping(self): + sink = InMemorySink() + audit = AuditLogger(sink) + app = _make_app(audit) + client = TestClient(app, raise_server_exceptions=False) + + client.post("/push") + http_events = [e for e in sink.events if e.event_type == "http.request"] + self.assertEqual(http_events[0].resource.type, "push_source") + self.assertIn("WRITE_ONLINE", http_events[0].resource.actions) diff --git a/sdk/python/tests/unit/audit/test_mcp_audit_handler.py b/sdk/python/tests/unit/audit/test_mcp_audit_handler.py new file mode 100644 index 00000000000..6a647e41e86 --- /dev/null +++ b/sdk/python/tests/unit/audit/test_mcp_audit_handler.py @@ -0,0 +1,382 @@ +"""Tests for MCP protocol-layer audit logging (handler wrapping). + +Validates ``_wrap_call_tool_handler`` and ``_principal_from_mcp_context`` +from ``feast.infra.mcp_servers.mcp_server``, using real ``FastApiMCP`` +and ``mcp.types.CallToolRequest`` objects so the tests break when the +upstream library changes its handler dispatch conventions. +""" + +import asyncio +import unittest +from types import SimpleNamespace +from typing import Any +from unittest.mock import MagicMock + +from fastapi import FastAPI + +from feast.audit.audit_logger import ( + AuditEvent, + AuditLogger, + AuditSink, + mcp_audit_request_id, +) +from feast.infra.mcp_servers.mcp_server import ( + _principal_from_mcp_context, + _wrap_call_tool_handler, +) + + +class InMemorySink(AuditSink): + def __init__(self): + self.events: list[AuditEvent] = [] + + def emit(self, event: AuditEvent) -> None: + self.events.append(event) + + +def _run(coro): + """Helper to run an async function synchronously in tests.""" + return asyncio.get_event_loop().run_until_complete(coro) + + +def _make_real_mcp(app: FastAPI | None = None): + """Build a *real* ``FastApiMCP`` instance backed by a throwaway FastAPI app. + + The returned object has ``server.request_handlers[CallToolRequest]`` + populated by fastapi_mcp, mirroring the production code path. + """ + from fastapi_mcp import FastApiMCP + + if app is None: + app = FastAPI() + + @app.get("/health") + async def health(): + return {"status": "ok"} + + mcp = FastApiMCP(app, name="test-feast", description="test") + mcp.mount() + return mcp + + +def _handler_key(): + from mcp.types import CallToolRequest + + return CallToolRequest + + +def _get_handler(mcp: Any): + """Return the current CallToolRequest handler from the real mcp server.""" + return mcp.server.request_handlers[_handler_key()] + + +def _make_call_tool_request( + tool_name: str = "some_tool", arguments: dict | None = None +): + """Create a real ``CallToolRequest`` pydantic object.""" + from mcp.types import CallToolRequest, CallToolRequestParams + + return CallToolRequest( + method="tools/call", + params=CallToolRequestParams(name=tool_name, arguments=arguments or {}), + ) + + +class TestWrapCallToolHandler(unittest.TestCase): + """Tests that exercise ``_wrap_call_tool_handler`` against a real + ``FastApiMCP.server.request_handlers`` dict. + """ + + def test_successful_call_logs_success(self): + sink = InMemorySink() + audit = AuditLogger(sink) + + mcp = _make_real_mcp() + original = _get_handler(mcp) + _wrap_call_tool_handler(mcp, audit) + + self.assertIsNot(_get_handler(mcp), original, "handler should be replaced") + + req = _make_call_tool_request("get_online_features") + try: + _run(_get_handler(mcp)(req)) + except Exception: + pass + + self.assertEqual(len(sink.events), 1) + event = sink.events[0] + self.assertEqual(event.event_type, "mcp.tools.call") + self.assertEqual(event.action.mcp_tool, "get_online_features") + self.assertEqual(event.source.transport, "mcp-http") + self.assertIsNotNone(event.duration_ms) + self.assertGreaterEqual(event.duration_ms, 0) + + def test_handler_exception_logs_error(self): + sink = InMemorySink() + audit = AuditLogger(sink) + + mcp = _make_real_mcp() + + async def exploding_handler(req): + raise RuntimeError("tool exploded") + + mcp.server.request_handlers[_handler_key()] = exploding_handler + _wrap_call_tool_handler(mcp, audit) + + req = _make_call_tool_request("failing_tool") + with self.assertRaises(RuntimeError): + _run(_get_handler(mcp)(req)) + + self.assertEqual(len(sink.events), 1) + event = sink.events[0] + self.assertEqual(event.outcome, "error") + self.assertEqual(event.action.mcp_tool, "failing_tool") + self.assertIn("tool exploded", event.detail) + + def test_isError_result_logs_mcp_error(self): + from mcp.types import CallToolResult, ServerResult, TextContent + + sink = InMemorySink() + audit = AuditLogger(sink) + + mcp = _make_real_mcp() + + async def error_handler(req): + return ServerResult( + root=CallToolResult( + content=[TextContent(type="text", text="validation failed")], + isError=True, + ) + ) + + mcp.server.request_handlers[_handler_key()] = error_handler + _wrap_call_tool_handler(mcp, audit) + + req = _make_call_tool_request("bad_tool") + _run(_get_handler(mcp)(req)) + + self.assertEqual(len(sink.events), 1) + self.assertEqual(sink.events[0].outcome, "mcp_error") + + def test_result_without_isError_attr_logs_success(self): + sink = InMemorySink() + audit = AuditLogger(sink) + + mcp = _make_real_mcp() + + async def plain_handler(req): + return {"content": "plain dict result"} + + mcp.server.request_handlers[_handler_key()] = plain_handler + _wrap_call_tool_handler(mcp, audit) + + req = _make_call_tool_request("simple_tool") + _run(_get_handler(mcp)(req)) + + self.assertEqual(sink.events[0].outcome, "success") + + def test_no_handler_skips_wrapping(self): + sink = InMemorySink() + audit = AuditLogger(sink) + + mcp = _make_real_mcp() + del mcp.server.request_handlers[_handler_key()] + + _wrap_call_tool_handler(mcp, audit) + self.assertNotIn(_handler_key(), mcp.server.request_handlers) + + def test_no_request_handlers_attr_is_safe(self): + sink = InMemorySink() + audit = AuditLogger(sink) + + mcp = SimpleNamespace(server=SimpleNamespace()) + _wrap_call_tool_handler(mcp, audit) + + def test_params_none_uses_empty_tool_name(self): + """When req.params is None the tool name defaults to empty string.""" + from mcp.types import CallToolResult, ServerResult + + sink = InMemorySink() + audit = AuditLogger(sink) + + mcp = _make_real_mcp() + + async def handler(req): + return ServerResult(root=CallToolResult(content=[])) + + mcp.server.request_handlers[_handler_key()] = handler + _wrap_call_tool_handler(mcp, audit) + + req = SimpleNamespace(params=None) + _run(_get_handler(mcp)(req)) + + self.assertEqual(sink.events[0].action.mcp_tool, "") + + def test_original_handler_result_is_returned(self): + sink = InMemorySink() + audit = AuditLogger(sink) + sentinel = object() + + mcp = _make_real_mcp() + + async def handler(req): + return sentinel + + mcp.server.request_handlers[_handler_key()] = handler + _wrap_call_tool_handler(mcp, audit) + + req = _make_call_tool_request("tool") + result = _run(_get_handler(mcp)(req)) + self.assertIs(result, sentinel) + + def test_contextvar_propagates_request_id(self): + """The wrapper sets mcp_audit_request_id so that AuditLoggingMiddleware + on the internal REST call can reuse the same request_id.""" + from mcp.types import CallToolResult, ServerResult + + sink = InMemorySink() + audit = AuditLogger(sink) + captured_ids: list[str | None] = [] + + mcp = _make_real_mcp() + + async def handler(req): + captured_ids.append(mcp_audit_request_id.get()) + return ServerResult(root=CallToolResult(content=[])) + + mcp.server.request_handlers[_handler_key()] = handler + _wrap_call_tool_handler(mcp, audit) + + req = _make_call_tool_request("tool") + _run(_get_handler(mcp)(req)) + + self.assertEqual(len(captured_ids), 1) + self.assertIsNotNone(captured_ids[0]) + self.assertEqual(captured_ids[0], sink.events[0].request_id) + + def test_contextvar_reset_after_call(self): + """The ContextVar is cleaned up after the handler completes.""" + from mcp.types import CallToolResult, ServerResult + + sink = InMemorySink() + audit = AuditLogger(sink) + + mcp = _make_real_mcp() + + async def handler(req): + return ServerResult(root=CallToolResult(content=[])) + + mcp.server.request_handlers[_handler_key()] = handler + _wrap_call_tool_handler(mcp, audit) + + self.assertIsNone(mcp_audit_request_id.get()) + + req = _make_call_tool_request("tool") + _run(_get_handler(mcp)(req)) + + self.assertIsNone(mcp_audit_request_id.get()) + + def test_mcp_and_rest_events_share_request_id(self): + """End-to-end: both mcp.tools.call and http.request events emitted + during a single tool invocation share the same request_id.""" + from mcp.types import CallToolResult, ServerResult + + sink = InMemorySink() + audit = AuditLogger(sink) + + mcp = _make_real_mcp() + + async def handler(req): + propagated = mcp_audit_request_id.get() + audit.log_http_request( + request_id=propagated or audit.new_request_id(), + method="POST", + path="/get-online-features", + status_code=200, + ) + return ServerResult(root=CallToolResult(content=[])) + + mcp.server.request_handlers[_handler_key()] = handler + _wrap_call_tool_handler(mcp, audit) + + req = _make_call_tool_request("get_online_features") + _run(_get_handler(mcp)(req)) + + self.assertEqual(len(sink.events), 2) + http_event = next(e for e in sink.events if e.event_type == "http.request") + mcp_event = next(e for e in sink.events if e.event_type == "mcp.tools.call") + self.assertEqual(http_event.request_id, mcp_event.request_id) + + def test_jsonrpc_id_from_request_context(self): + """When request_context is available, jsonrpc_id is captured.""" + from mcp.types import CallToolResult, ServerResult + + sink = InMemorySink() + audit = AuditLogger(sink) + + mcp = _make_real_mcp() + + async def handler(req): + return ServerResult(root=CallToolResult(content=[])) + + mcp.server.request_handlers[_handler_key()] = handler + _wrap_call_tool_handler(mcp, audit) + + req = _make_call_tool_request("some_tool") + _run(_get_handler(mcp)(req)) + + # Without a live transport, request_context raises LookupError, + # so jsonrpc_id should be None. + self.assertIsNone(sink.events[0].jsonrpc_id) + + +class TestPrincipalFromMcpContext(unittest.TestCase): + def test_extracts_auth_type_header(self): + request = MagicMock() + request.headers = {"x-feast-auth-type": "oidc", "authorization": "Bearer tok"} + ctx = SimpleNamespace(request=request) + + server = SimpleNamespace(request_context=ctx) + principal = _principal_from_mcp_context(server) + self.assertEqual(principal.auth_type, "oidc") + self.assertEqual(principal.username, "(authenticated)") + + def test_no_auth_header_returns_empty_username(self): + request = MagicMock() + request.headers = {} + ctx = SimpleNamespace(request=request) + + server = SimpleNamespace(request_context=ctx) + principal = _principal_from_mcp_context(server) + self.assertEqual(principal.username, "") + self.assertEqual(principal.auth_type, "") + + def test_no_request_returns_empty_principal(self): + ctx = SimpleNamespace() + server = SimpleNamespace(request_context=ctx) + principal = _principal_from_mcp_context(server) + self.assertEqual(principal.username, "") + + def test_lookup_error_returns_empty_principal(self): + """When request_context ContextVar is not set, LookupError is raised.""" + + class FakeServer: + @property + def request_context(self): + raise LookupError("no context") + + principal = _principal_from_mcp_context(FakeServer()) + self.assertEqual(principal.username, "") + + def test_none_server_returns_empty_principal(self): + principal = _principal_from_mcp_context(None) + self.assertEqual(principal.username, "") + + def test_exception_in_headers_returns_empty_principal(self): + request = MagicMock() + request.headers = property(lambda self: (_ for _ in ()).throw(RuntimeError)) + ctx = SimpleNamespace(request=request) + + server = SimpleNamespace(request_context=ctx) + principal = _principal_from_mcp_context(server) + self.assertEqual(principal.username, "")