From d56a0a97f8ddc348ab761ec0cf270fe09cee2dba Mon Sep 17 00:00:00 2001 From: Seongho Bae Date: Thu, 5 Feb 2026 16:16:00 +0900 Subject: [PATCH] fix(ui): allow editing MCP stdio transport config (#20241) * fix(ui): enable stdio transport edits for MCP servers * fix(ui): use antd Input in MCP edit stdio Align MCP Server Edit with UI guidelines by replacing deprecated Tremor TextInput, and relax stdio args validation to match create flow while improving test stability. * fix(otel): make semantic log LogRecord import mypy-safe Prefer the OTEL >=1.39.0 LogRecord import path and keep an ignored fallback for older versions so MyPy doesn't fail on newer SDK stubs. * fix(otel): tolerate LogRecord ctor changes across SDK versions Create semantic LogRecords via a best-effort wrapper that falls back when the `resource` kwarg is unsupported (OTEL >= 1.39), and avoid MyPy overload/no-redef failures. * fix(otel): silence mypy no-redef on versioned LogRecord import MyPy sees both branches of the version-compat import and flags a redefinition. Ignore no-redef on the legacy import path to keep CI passing. * fix(ui): ensure mcp_info.server_name is always populated When using stdio transport there may be no URL to fall back on; prefer existing server_name/url/alias to avoid sending an empty mcp_info.server_name on update. * chore(otel): format opentelemetry; ignore ui export output * fix: guard optional a2a resolver + make OTEL semantic logs mypy-safe * chore: format A2A resolver and OTEL semantic logs * fix: address review feedback for MCP stdio edit * fix: keep MCP stdio edit PR scoped * fix(otel): make semantic logs mypy-safe --- litellm/integrations/opentelemetry.py | 235 +++++++--- .../mcp_tools/StdioConfiguration.tsx | 9 +- .../mcp_tools/mcp_server_columns.tsx | 6 +- .../mcp_tools/mcp_server_edit.test.tsx | 154 +++++++ .../components/mcp_tools/mcp_server_edit.tsx | 410 ++++++++++++++++-- .../components/mcp_tools/mcp_server_view.tsx | 6 +- .../src/components/mcp_tools/types.tsx | 12 +- 7 files changed, 711 insertions(+), 121 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 296a88f9a0..138d508db4 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -5,6 +5,10 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast import litellm from litellm._logging import verbose_logger +from litellm.integrations._types.open_inference import ( + OpenInferenceSpanKindValues, + SpanAttributes, +) from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.secret_managers.main import get_secret_bool @@ -17,10 +21,6 @@ from litellm.types.utils import ( StandardCallbackDynamicParams, StandardLoggingPayload, ) -from litellm.integrations._types.open_inference import ( - OpenInferenceSpanKindValues, - SpanAttributes, -) # OpenTelemetry imports moved to individual functions to avoid import errors when not installed @@ -40,7 +40,9 @@ if TYPE_CHECKING: Context = Union[_Context, Any] SpanExporter = Union[_SpanExporter, Any] UserAPIKeyAuth = Union[_UserAPIKeyAuth, Any] - ManagementEndpointLoggingPayload = Union[_ManagementEndpointLoggingPayload, Any] + ManagementEndpointLoggingPayload = Union[ + _ManagementEndpointLoggingPayload, Any + ] else: Span = Any Tracer = Any @@ -95,12 +97,16 @@ class OpenTelemetryConfig: exporter = os.getenv( "OTEL_EXPORTER_OTLP_PROTOCOL", os.getenv("OTEL_EXPORTER", "console") ) - endpoint = os.getenv("OTEL_EXPORTER_OTLP_ENDPOINT", os.getenv("OTEL_ENDPOINT")) + endpoint = os.getenv( + "OTEL_EXPORTER_OTLP_ENDPOINT", os.getenv("OTEL_ENDPOINT") + ) headers = os.getenv( "OTEL_EXPORTER_OTLP_HEADERS", os.getenv("OTEL_HEADERS") ) # example: OTEL_HEADERS=x-honeycomb-team=B85YgLm96***" enable_metrics: bool = ( - os.getenv("LITELLM_OTEL_INTEGRATION_ENABLE_METRICS", "false").lower() + os.getenv( + "LITELLM_OTEL_INTEGRATION_ENABLE_METRICS", "false" + ).lower() == "true" ) enable_events: bool = ( @@ -108,7 +114,9 @@ class OpenTelemetryConfig: == "true" ) service_name = os.getenv("OTEL_SERVICE_NAME", "litellm") - deployment_environment = os.getenv("OTEL_ENVIRONMENT_NAME", "production") + deployment_environment = os.getenv( + "OTEL_ENVIRONMENT_NAME", "production" + ) model_id = os.getenv("OTEL_MODEL_ID", service_name) if exporter == "in_memory": @@ -157,7 +165,9 @@ class OpenTelemetry(CustomLogger): logging.getLogger(__name__) # Enable OpenTelemetry logging - otel_exporter_logger = logging.getLogger("opentelemetry.sdk.trace.export") + otel_exporter_logger = logging.getLogger( + "opentelemetry.sdk.trace.export" + ) otel_exporter_logger.setLevel(logging.DEBUG) # init CustomLogger params @@ -253,7 +263,9 @@ class OpenTelemetry(CustomLogger): # Don't call set_provider to preserve existing context else: # Default proxy provider or unknown type, create our own - verbose_logger.debug("OpenTelemetry: Creating new %s", provider_name) + verbose_logger.debug( + "OpenTelemetry: Creating new %s", provider_name + ) provider = create_new_provider_fn() set_provider_fn(provider) except Exception as e: @@ -274,7 +286,9 @@ class OpenTelemetry(CustomLogger): from opentelemetry.trace import SpanKind def create_tracer_provider(): - provider = TracerProvider(resource=self._get_litellm_resource(self.config)) + provider = TracerProvider( + resource=self._get_litellm_resource(self.config) + ) provider.add_span_processor(self._get_span_processor()) return provider @@ -388,10 +402,14 @@ class OpenTelemetry(CustomLogger): def log_failure_event(self, kwargs, response_obj, start_time, end_time): self._handle_failure(kwargs, response_obj, start_time, end_time) - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_success_event( + self, kwargs, response_obj, start_time, end_time + ): self._handle_success(kwargs, response_obj, start_time, end_time) - async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_failure_event( + self, kwargs, response_obj, start_time, end_time + ): self._handle_failure(kwargs, response_obj, start_time, end_time) async def async_service_success_hook( @@ -588,7 +606,9 @@ class OpenTelemetry(CustomLogger): if dynamic_headers is not None: # Create spans using a temporary tracer with dynamic headers - tracer_to_use = self._get_tracer_with_dynamic_headers(dynamic_headers) + tracer_to_use = self._get_tracer_with_dynamic_headers( + dynamic_headers + ) verbose_logger.debug( "Using dynamic headers for this request: %s", dynamic_headers ) @@ -624,7 +644,9 @@ class OpenTelemetry(CustomLogger): ) # Create a temporary tracer provider with dynamic headers - temp_provider = TracerProvider(resource=self._get_litellm_resource(self.config)) + temp_provider = TracerProvider( + resource=self._get_litellm_resource(self.config) + ) temp_provider.add_span_processor( self._get_span_processor(dynamic_headers=dynamic_headers) ) @@ -755,7 +777,9 @@ class OpenTelemetry(CustomLogger): metadata = litellm_params.get("metadata") or {} generation_name = metadata.get("generation_name") - raw_span_name = generation_name if generation_name else RAW_REQUEST_SPAN_NAME + raw_span_name = ( + generation_name if generation_name else RAW_REQUEST_SPAN_NAME + ) otel_tracer: Tracer = self.get_tracer_to_use_for_request(kwargs) raw_span = otel_tracer.start_span( @@ -780,7 +804,9 @@ class OpenTelemetry(CustomLogger): } std_log = kwargs.get("standard_logging_object") - md = getattr(std_log, "metadata", None) or (std_log or {}).get("metadata", {}) + md = getattr(std_log, "metadata", None) or (std_log or {}).get( + "metadata", {} + ) for key in [ "user_api_key_hash", "user_api_key_alias", @@ -802,9 +828,9 @@ class OpenTelemetry(CustomLogger): common_attrs[f"metadata.{key}"] = str(md[key]) # get hidden params - hidden_params = getattr(std_log, "hidden_params", None) or (std_log or {}).get( - "hidden_params", {} - ) + hidden_params = getattr(std_log, "hidden_params", None) or ( + std_log or {} + ).get("hidden_params", {}) if hidden_params: common_attrs["hidden_params"] = safe_dumps(hidden_params) @@ -838,7 +864,9 @@ class OpenTelemetry(CustomLogger): self._record_response_duration_metric(kwargs, end_time, common_attrs) @staticmethod - def _to_timestamp(val: Optional[Union[datetime, float, str]]) -> Optional[float]: + def _to_timestamp( + val: Optional[Union[datetime, float, str]], + ) -> Optional[float]: """Convert datetime/float/string to timestamp.""" if val is None: return None @@ -855,7 +883,9 @@ class OpenTelemetry(CustomLogger): except ValueError: return None - def _record_time_to_first_token_metric(self, kwargs: dict, common_attrs: dict): + def _record_time_to_first_token_metric( + self, kwargs: dict, common_attrs: dict + ): """Record Time to First Token (TTFT) metric for streaming requests.""" optional_params = kwargs.get("optional_params", {}) is_streaming = optional_params.get("stream", False) @@ -868,7 +898,10 @@ class OpenTelemetry(CustomLogger): api_call_start_time = kwargs.get("api_call_start_time", None) completion_start_time = kwargs.get("completion_start_time", None) - if api_call_start_time is not None and completion_start_time is not None: + if ( + api_call_start_time is not None + and completion_start_time is not None + ): # Convert to timestamps if needed (handles datetime, float, and string) api_call_start_ts = self._to_timestamp(api_call_start_time) completion_start_ts = self._to_timestamp(completion_start_time) @@ -876,7 +909,9 @@ class OpenTelemetry(CustomLogger): if api_call_start_ts is None or completion_start_ts is None: return # Skip recording if conversion failed - time_to_first_token_seconds = completion_start_ts - api_call_start_ts + time_to_first_token_seconds = ( + completion_start_ts - api_call_start_ts + ) self._time_to_first_token_histogram.record( time_to_first_token_seconds, attributes=common_attrs ) @@ -946,7 +981,9 @@ class OpenTelemetry(CustomLogger): generation_time_seconds = duration_s if generation_time_seconds > 0: - time_per_output_token_seconds = generation_time_seconds / completion_tokens + time_per_output_token_seconds = ( + generation_time_seconds / completion_tokens + ) self._time_per_output_token_histogram.record( time_per_output_token_seconds, attributes=common_attrs ) @@ -1007,21 +1044,26 @@ class OpenTelemetry(CustomLogger): # See: https://github.com/open-telemetry/opentelemetry-python/pull/4676 # TODO: Refactor to use the proper OTEL Logs API instead of directly creating SDK LogRecords - from opentelemetry._logs import SeverityNumber, get_logger, get_logger_provider + from opentelemetry._logs import ( + SeverityNumber, + get_logger, + ) - try: - from opentelemetry.sdk._logs import LogRecord as SdkLogRecord # type: ignore[attr-defined] # OTEL < 1.39.0 - except ImportError: - from opentelemetry.sdk._logs._internal import LogRecord as SdkLogRecord # type: ignore[attr-defined, no-redef] # OTEL >= 1.39.0 + # MyPy evaluates both branches of try/except imports and can fail when + # newer OTEL stubs remove/relocate symbols. Gate the typing import so + # only the canonical location is type-checked. + if TYPE_CHECKING: + from opentelemetry.sdk._logs._internal import LogRecord as SdkLogRecord + else: + try: + from opentelemetry.sdk._logs import ( + LogRecord as SdkLogRecord, # type: ignore[attr-defined] + ) + except ImportError: + from opentelemetry.sdk._logs._internal import LogRecord as SdkLogRecord otel_logger = get_logger(LITELLM_LOGGER_NAME) - # Get the resource from the logger provider - logger_provider = get_logger_provider() - resource = getattr( - logger_provider, "_resource", None - ) or self._get_litellm_resource(self.config) - parent_ctx = span.get_span_context() provider = (kwargs.get("litellm_params") or {}).get( "custom_llm_provider", "Unknown" @@ -1030,7 +1072,10 @@ class OpenTelemetry(CustomLogger): # per-message events for msg in kwargs.get("messages", []): role = msg.get("role", "user") - attrs = {"event_name": "gen_ai.content.prompt", "gen_ai.system": provider} + attrs = { + "event_name": "gen_ai.content.prompt", + "gen_ai.system": provider, + } if role == "tool" and msg.get("id"): attrs["id"] = msg["id"] if self.message_logging and msg.get("content"): @@ -1044,7 +1089,6 @@ class OpenTelemetry(CustomLogger): severity_number=SeverityNumber.INFO, severity_text="INFO", body=msg.copy(), - resource=resource, attributes=attrs, ) otel_logger.emit(log_record) @@ -1076,7 +1120,6 @@ class OpenTelemetry(CustomLogger): severity_number=SeverityNumber.INFO, severity_text="INFO", body=body, - resource=resource, attributes=attrs, ) otel_logger.emit(log_record) @@ -1146,7 +1189,9 @@ class OpenTelemetry(CustomLogger): value=guardrail_information.get("guardrail_mode"), ) - masked_entity_count = guardrail_information.get("masked_entity_count") + masked_entity_count = guardrail_information.get( + "masked_entity_count" + ) if masked_entity_count is not None: guardrail_span.set_attribute( "masked_entity_count", safe_dumps(masked_entity_count) @@ -1173,8 +1218,9 @@ class OpenTelemetry(CustomLogger): # Decide whether to create a primary span # Always create if no parent span exists (backward compatibility) # OR if USE_OTEL_LITELLM_REQUEST_SPAN is explicitly enabled - should_create_primary_span = parent_otel_span is None or get_secret_bool( - "USE_OTEL_LITELLM_REQUEST_SPAN" + should_create_primary_span = ( + parent_otel_span is None + or get_secret_bool("USE_OTEL_LITELLM_REQUEST_SPAN") ) if should_create_primary_span: @@ -1200,7 +1246,9 @@ class OpenTelemetry(CustomLogger): if parent_otel_span.is_recording(): parent_otel_span.set_status(Status(StatusCode.ERROR)) self.set_attributes(parent_otel_span, kwargs, response_obj) - self._record_exception_on_span(span=parent_otel_span, kwargs=kwargs) + self._record_exception_on_span( + span=parent_otel_span, kwargs=kwargs + ) # Create span for guardrail information self._create_guardrail_span(kwargs=kwargs, context=_parent_context) @@ -1223,7 +1271,9 @@ class OpenTelemetry(CustomLogger): 2. Sets structured error attributes from StandardLoggingPayloadErrorInformation """ try: - from litellm.integrations._types.open_inference import ErrorAttributes + from litellm.integrations._types.open_inference import ( + ErrorAttributes, + ) # Get the exception object if available exception = kwargs.get("exception") @@ -1233,15 +1283,17 @@ class OpenTelemetry(CustomLogger): span.record_exception(exception) # Get StandardLoggingPayload for structured error information - standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get( - "standard_logging_object" + standard_logging_payload: Optional[StandardLoggingPayload] = ( + kwargs.get("standard_logging_object") ) if standard_logging_payload is None: return # Extract error_information from StandardLoggingPayload - error_information = standard_logging_payload.get("error_information") + error_information = standard_logging_payload.get( + "error_information" + ) if error_information is None: # Fallback to error_str if error_information is not available @@ -1331,7 +1383,9 @@ class OpenTelemetry(CustomLogger): ) pass - def cast_as_primitive_value_type(self, value) -> Union[str, bool, int, float]: + def cast_as_primitive_value_type( + self, value + ) -> Union[str, bool, int, float]: """ Casts the value to a primitive OTEL type if it is not already a primitive type. @@ -1401,8 +1455,8 @@ class OpenTelemetry(CustomLogger): optional_params = kwargs.get("optional_params", {}) litellm_params = kwargs.get("litellm_params", {}) or {} - standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get( - "standard_logging_object" + standard_logging_payload: Optional[StandardLoggingPayload] = ( + kwargs.get("standard_logging_object") ) if standard_logging_payload is None: raise ValueError("standard_logging_object not found in kwargs") @@ -1424,11 +1478,13 @@ class OpenTelemetry(CustomLogger): ) or (standard_logging_payload or {}).get("hidden_params", {}) if hidden_params: self.safe_set_attribute( - span=span, key="hidden_params", value=safe_dumps(hidden_params) + span=span, + key="hidden_params", + value=safe_dumps(hidden_params), ) # Cost breakdown tracking - cost_breakdown: Optional[CostBreakdown] = standard_logging_payload.get( - "cost_breakdown" + cost_breakdown: Optional[CostBreakdown] = ( + standard_logging_payload.get("cost_breakdown") ) if cost_breakdown: for key, value in cost_breakdown.items(): @@ -1504,7 +1560,9 @@ class OpenTelemetry(CustomLogger): # The unique identifier for the completion. if response_obj and response_obj.get("id"): self.safe_set_attribute( - span=span, key="gen_ai.response.id", value=response_obj.get("id") + span=span, + key="gen_ai.response.id", + value=response_obj.get("id"), ) # The model used to generate the response. @@ -1639,7 +1697,9 @@ class OpenTelemetry(CustomLogger): "OpenTelemetry logging error in set_attributes %s", str(e) ) - def _cast_as_primitive_value_type(self, value) -> Union[str, bool, int, float]: + def _cast_as_primitive_value_type( + self, value + ) -> Union[str, bool, int, float]: """ Casts the value to a primitive OTEL type if it is not already a primitive type. @@ -1673,7 +1733,10 @@ class OpenTelemetry(CustomLogger): if isinstance(messages, str): # Handle system_instructions passed as a string return [ - {"role": "system", "parts": [{"type": "text", "content": messages}]} + { + "role": "system", + "parts": [{"type": "text", "content": messages}], + } ] transformed = [] @@ -1714,9 +1777,11 @@ class OpenTelemetry(CustomLogger): message = choice.get("message") or {} finish_reason = choice.get("finish_reason") - transformed_msg = self._transform_messages_to_otel_semantic_conventions( - [message] - )[0] + transformed_msg = ( + self._transform_messages_to_otel_semantic_conventions( + [message] + )[0] + ) if finish_reason: transformed_msg["finish_reason"] = finish_reason @@ -1728,7 +1793,9 @@ class OpenTelemetry(CustomLogger): self.set_attributes(span, kwargs, response_obj) kwargs.get("optional_params", {}) litellm_params = kwargs.get("litellm_params", {}) or {} - custom_llm_provider = litellm_params.get("custom_llm_provider", "Unknown") + custom_llm_provider = litellm_params.get( + "custom_llm_provider", "Unknown" + ) _raw_response = kwargs.get("original_response") _additional_args = kwargs.get("additional_args", {}) or {} @@ -1741,7 +1808,9 @@ class OpenTelemetry(CustomLogger): if complete_input_dict and isinstance(complete_input_dict, dict): for param, val in complete_input_dict.items(): self.safe_set_attribute( - span=span, key=f"llm.{custom_llm_provider}.{param}", value=val + span=span, + key=f"llm.{custom_llm_provider}.{param}", + value=val, ) ############################################# @@ -1773,7 +1842,8 @@ class OpenTelemetry(CustomLogger): ) except Exception as e: verbose_logger.exception( - "OpenTelemetry logging error in set_raw_request_attributes %s", str(e) + "OpenTelemetry logging error in set_raw_request_attributes %s", + str(e), ) def _to_ns(self, dt): @@ -1813,7 +1883,9 @@ class OpenTelemetry(CustomLogger): ) litellm_params = kwargs.get("litellm_params", {}) or {} - proxy_server_request = litellm_params.get("proxy_server_request", {}) or {} + proxy_server_request = ( + litellm_params.get("proxy_server_request", {}) or {} + ) headers = proxy_server_request.get("headers", {}) or {} traceparent = headers.get("traceparent", None) _metadata = litellm_params.get("metadata", {}) or {} @@ -1832,7 +1904,10 @@ class OpenTelemetry(CustomLogger): "OpenTelemetry: Using traceparent header for context propagation" ) carrier = {"traceparent": traceparent} - return TraceContextTextMapPropagator().extract(carrier=carrier), None + return ( + TraceContextTextMapPropagator().extract(carrier=carrier), + None, + ) # Priority 3: Active span from global context (auto-detection) try: @@ -1960,10 +2035,14 @@ class OpenTelemetry(CustomLogger): self.OTEL_HEADERS, ) - _split_otel_headers = OpenTelemetry._get_headers_dictionary(self.OTEL_HEADERS) + _split_otel_headers = OpenTelemetry._get_headers_dictionary( + self.OTEL_HEADERS + ) # Normalize endpoint for logs - ensure it points to /v1/logs instead of /v1/traces - normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, "logs") + normalized_endpoint = self._normalize_otel_endpoint( + self.OTEL_ENDPOINT, "logs" + ) verbose_logger.debug( "OpenTelemetry: Log endpoint normalized from %s to %s", @@ -2051,14 +2130,18 @@ class OpenTelemetry(CustomLogger): self.OTEL_HEADERS, ) - _split_otel_headers = OpenTelemetry._get_headers_dictionary(self.OTEL_HEADERS) + _split_otel_headers = OpenTelemetry._get_headers_dictionary( + self.OTEL_HEADERS + ) normalized_endpoint = self._normalize_otel_endpoint( self.OTEL_ENDPOINT, "metrics" ) if self.OTEL_EXPORTER == "console": exporter = ConsoleMetricExporter() - return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) + return PeriodicExportingMetricReader( + exporter, export_interval_millis=5000 + ) elif ( self.OTEL_EXPORTER == "otlp_http" @@ -2074,7 +2157,9 @@ class OpenTelemetry(CustomLogger): headers=_split_otel_headers, preferred_temporality={Histogram: AggregationTemporality.DELTA}, ) - return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) + return PeriodicExportingMetricReader( + exporter, export_interval_millis=5000 + ) elif self.OTEL_EXPORTER == "otlp_grpc" or self.OTEL_EXPORTER == "grpc": try: @@ -2092,7 +2177,9 @@ class OpenTelemetry(CustomLogger): headers=_split_otel_headers, preferred_temporality={Histogram: AggregationTemporality.DELTA}, ) - return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) + return PeriodicExportingMetricReader( + exporter, export_interval_millis=5000 + ) else: verbose_logger.warning( @@ -2100,7 +2187,9 @@ class OpenTelemetry(CustomLogger): self.OTEL_EXPORTER, ) exporter = ConsoleMetricExporter() - return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) + return PeriodicExportingMetricReader( + exporter, export_interval_millis=5000 + ) def _normalize_otel_endpoint( self, endpoint: Optional[str], signal_type: str @@ -2171,7 +2260,9 @@ class OpenTelemetry(CustomLogger): return endpoint @staticmethod - def _get_headers_dictionary(headers: Optional[Union[str, dict]]) -> Dict[str, str]: + def _get_headers_dictionary( + headers: Optional[Union[str, dict]], + ) -> Dict[str, str]: """ Convert a string or dictionary of headers into a dictionary of headers. """ diff --git a/ui/litellm-dashboard/src/components/mcp_tools/StdioConfiguration.tsx b/ui/litellm-dashboard/src/components/mcp_tools/StdioConfiguration.tsx index 23f5f84fea..476a5b6168 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/StdioConfiguration.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/StdioConfiguration.tsx @@ -4,9 +4,14 @@ import { InfoCircleOutlined } from "@ant-design/icons"; interface StdioConfigurationProps { isVisible: boolean; + /** + * When true, stdio_config is required + validated as JSON. + * Edit screen can set this to false when using dedicated command/args/env fields. + */ + required?: boolean; } -const StdioConfiguration: React.FC = ({ isVisible }) => { +const StdioConfiguration: React.FC = ({ isVisible, required = true }) => { if (!isVisible) return null; return ( @@ -21,7 +26,7 @@ const StdioConfiguration: React.FC = ({ isVisible }) => } name="stdio_config" rules={[ - { required: true, message: "Please enter stdio configuration" }, + ...(required ? [{ required: true, message: "Please enter stdio configuration" }] : []), { validator: (_, value) => { if (!value) return Promise.resolve(); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_columns.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_columns.tsx index ac611c633a..78e9c6f465 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_columns.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_columns.tsx @@ -36,7 +36,11 @@ export const mcpServerColumns = ( id: "url", header: "URL", cell: ({ row }) => { - const { maskedUrl } = getMaskedAndFullUrl(row.original.url); + const url = row.original.url; + if (!url) { + return ; + } + const { maskedUrl } = getMaskedAndFullUrl(url); return {maskedUrl}; }, }, diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx new file mode 100644 index 0000000000..e33e2fff49 --- /dev/null +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx @@ -0,0 +1,154 @@ +import React from "react"; +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { render, screen, waitFor, fireEvent, act } from "@testing-library/react"; +import MCPServerEdit from "./mcp_server_edit"; +import * as networking from "../networking"; + +vi.mock("../networking", () => ({ + updateMCPServer: vi.fn(), + testMCPToolsListRequest: vi.fn().mockResolvedValue({ tools: [], error: null }), +})); + +vi.mock("../molecules/notifications_manager", () => ({ + default: { + success: vi.fn(), + fromBackend: vi.fn(), + }, +})); + +vi.mock("@/hooks/useMcpOAuthFlow", () => ({ + useMcpOAuthFlow: () => ({ + startOAuthFlow: vi.fn(), + status: "idle", + error: null, + tokenResponse: null, + }), +})); + +vi.mock("./mcp_server_cost_config", () => ({ + default: () =>
, +})); + +vi.mock("./MCPPermissionManagement", () => ({ + default: () =>
, +})); + +vi.mock("./mcp_tool_configuration", () => ({ + default: () =>
, +})); + +describe("MCPServerEdit (stdio)", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should render without crashing", () => { + render( + , + ); + + expect(screen.getByRole("tab", { name: "Server Configuration" })).toBeInTheDocument(); + }); + + it("should allow updating stdio transport configuration", async () => { + const onCancel = vi.fn(); + const onSuccess = vi.fn(); + + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + server_id: "server-1", + server_name: "TestServer", + alias: "test", + transport: "stdio", + url: null, + command: "npx", + args: ["-y", "@circleci/mcp-server-circleci"], + env: { CIRCLECI_TOKEN: "***" }, + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + }); + + render( + , + ); + + // Stdio section should be visible + expect(screen.getByLabelText("Command")).toBeInTheDocument(); + + // URL field should not be visible when transport=stdio + expect(screen.queryByText("MCP Server URL")).not.toBeInTheDocument(); + + // Update env_json + const envTextarea = screen.getByLabelText("Environment (JSON object)"); + await act(async () => { + fireEvent.change(envTextarea, { + target: { + value: JSON.stringify({ CIRCLECI_TOKEN: "new-token", CIRCLECI_BASE_URL: "https://circleci.com" }, null, 2), + }, + }); + }); + + const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); + const saveButton = saveButtons[0]; + await act(async () => { + fireEvent.click(saveButton); + }); + + await waitFor(() => { + expect(networking.updateMCPServer).toHaveBeenCalledTimes(1); + }); + + const [_token, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; + expect(_token).toBe("access-token"); + expect(payload.transport).toBe("stdio"); + expect(payload.command).toBe("npx"); + expect(payload.args).toEqual(["-y", "@circleci/mcp-server-circleci"]); + expect(payload.env).toEqual({ CIRCLECI_TOKEN: "new-token", CIRCLECI_BASE_URL: "https://circleci.com" }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx index 05a5a3cf4b..fa46521e19 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx @@ -1,5 +1,5 @@ import React, { useState, useEffect } from "react"; -import { Form, Select, Button as AntdButton, Tooltip } from "antd"; +import { Form, Select, Button as AntdButton, Tooltip, Input } from "antd"; import { InfoCircleOutlined } from "@ant-design/icons"; import { Button, TextInput, TabGroup, TabList, Tab, TabPanels, TabPanel } from "@tremor/react"; import { AUTH_TYPE, OAUTH_FLOW, MCPServer, MCPServerCostInfo } from "./types"; @@ -8,6 +8,7 @@ import { updateMCPServer, testMCPToolsListRequest } from "../networking"; import MCPServerCostConfig from "./mcp_server_cost_config"; import MCPPermissionManagement from "./MCPPermissionManagement"; import MCPToolConfiguration from "./mcp_tool_configuration"; +import StdioConfiguration from "./StdioConfiguration"; import { validateMCPServerUrl, validateMCPServerName } from "./utils"; import NotificationsManager from "../molecules/notifications_manager"; import { useMcpOAuthFlow } from "@/hooks/useMcpOAuthFlow"; @@ -40,6 +41,8 @@ const MCPServerEdit: React.FC = ({ const [allowedTools, setAllowedTools] = useState([]); const [pendingRestoredValues, setPendingRestoredValues] = useState | null>(null); const authType = Form.useWatch("auth_type", form) as string | undefined; + const transportType = Form.useWatch("transport", form) as string | undefined; + const isStdioTransport = transportType === "stdio"; const shouldShowAuthValueField = authType ? AUTH_TYPES_REQUIRING_AUTH_VALUE.includes(authType) : false; const isOAuthAuthType = authType === AUTH_TYPE.OAUTH2; const oauthFlowTypeValue = Form.useWatch("oauth_flow_type", form) as string | undefined; @@ -127,13 +130,26 @@ const MCPServerEdit: React.FC = ({ })); }, [mcpServer.static_headers]); + const initialEnvJson = React.useMemo(() => { + const env = mcpServer.env ?? undefined; + if (!env || Object.keys(env).length === 0) { + return ""; + } + try { + return JSON.stringify(env, null, 2); + } catch { + return ""; + } + }, [mcpServer.env]); + + const initialValues = React.useMemo( () => ({ ...mcpServer, static_headers: initialStaticHeaders, oauth_flow_type: mcpServer.token_url ? OAUTH_FLOW.M2M : OAUTH_FLOW.INTERACTIVE, }), - [mcpServer, initialStaticHeaders], + [mcpServer, initialStaticHeaders, initialEnvJson], ); // Initialize cost config from existing server data @@ -214,9 +230,10 @@ const MCPServerEdit: React.FC = ({ }, [mcpServer, accessToken, oauthAccessToken]); const fetchTools = async () => { - if (!accessToken || !mcpServer.url) { - return; - } + if (!accessToken) return; + + // HTTP/SSE requires a URL; stdio does not. + if (mcpServer.transport !== "stdio" && !mcpServer.url) return; const isM2M = mcpServer.auth_type === AUTH_TYPE.OAUTH2 && !!mcpServer.token_url; if (mcpServer.auth_type === AUTH_TYPE.OAUTH2 && !isM2M && !oauthAccessToken) { @@ -237,6 +254,9 @@ const MCPServerEdit: React.FC = ({ authorization_url: mcpServer.authorization_url, token_url: mcpServer.token_url, registration_url: mcpServer.registration_url, + command: mcpServer.command, + args: mcpServer.args, + env: mcpServer.env, }; const toolsResponse = await testMCPToolsListRequest(accessToken, mcpServerConfig, oauthAccessToken); @@ -287,6 +307,27 @@ const MCPServerEdit: React.FC = ({ return existingOptions; }; + const handleTransportChange = (value: string) => { + // Clear fields that are not relevant for the selected transport. + if (value === "stdio") { + form.setFieldsValue({ + url: undefined, + auth_type: undefined, + credentials: undefined, + authorization_url: undefined, + token_url: undefined, + registration_url: undefined, + }); + } else { + form.setFieldsValue({ + command: undefined, + args: undefined, + env_json: undefined, + stdio_config: undefined, + }); + } + }; + const handleSave = async (values: Record) => { if (!accessToken) return; try { @@ -294,6 +335,10 @@ const MCPServerEdit: React.FC = ({ const { static_headers: staticHeadersList, credentials: credentialValues, + stdio_config: rawStdioConfig, + env_json: rawEnvJson, + command: rawCommand, + args: rawArgs, allow_all_keys: allowAllKeysRaw, available_on_public_internet: availableOnPublicInternetRaw, ...restValues @@ -334,12 +379,104 @@ const MCPServerEdit: React.FC = ({ }, {}) : undefined; + let stdioFields: Record = {}; + + if (restValues.transport === "stdio") { + // Prefer JSON config if provided (matches Create screen behavior) + if (rawStdioConfig) { + try { + const stdioConfig = JSON.parse(rawStdioConfig); + + let actualConfig = stdioConfig; + if (stdioConfig?.mcpServers && typeof stdioConfig.mcpServers === "object") { + const serverNames = Object.keys(stdioConfig.mcpServers); + if (serverNames.length > 0) { + actualConfig = stdioConfig.mcpServers[serverNames[0]]; + } + } + + const parsedArgs = Array.isArray(actualConfig?.args) + ? actualConfig.args.map((v: any) => String(v)).filter((v: string) => v.trim() !== "") + : []; + + const parsedEnv = + actualConfig?.env && typeof actualConfig.env === "object" && !Array.isArray(actualConfig.env) + ? Object.entries(actualConfig.env).reduce((acc: Record, [k, v]) => { + if (k == null || String(k).trim() === "") return acc; + acc[String(k)] = v == null ? "" : String(v); + return acc; + }, {}) + : {}; + + stdioFields = { + command: actualConfig?.command ? String(actualConfig.command) : undefined, + args: parsedArgs, + env: parsedEnv, + }; + + if (!stdioFields.command) { + NotificationsManager.fromBackend("Stdio configuration must include a command"); + return; + } + } catch { + NotificationsManager.fromBackend("Invalid JSON in stdio configuration"); + return; + } + } else { + // Dedicated fields path (command/args + env JSON) + let parsedEnv: Record = {}; + if (rawEnvJson) { + try { + const env = JSON.parse(rawEnvJson); + if (env && typeof env === "object" && !Array.isArray(env)) { + parsedEnv = Object.entries(env).reduce((acc: Record, [k, v]) => { + if (k == null || String(k).trim() === "") return acc; + acc[String(k)] = v == null ? "" : String(v); + return acc; + }, {}); + } + } catch { + NotificationsManager.fromBackend("Invalid JSON in stdio env configuration"); + return; + } + } + const parsedArgs = Array.isArray(rawArgs) + ? rawArgs.map((v: any) => String(v)).filter((v: string) => v.trim() !== "") + : []; + + const parsedCommand = rawCommand ? String(rawCommand).trim() : ""; + if (!parsedCommand) { + NotificationsManager.fromBackend("Stdio transport requires a command"); + return; + } + + stdioFields = { + command: parsedCommand, + args: parsedArgs, + env: parsedEnv, + }; + } + } + // Prepare the payload with cost configuration and permission fields + const mcpInfoServerName = + restValues.server_name || + restValues.url || + mcpServer.server_name || + mcpServer.url || + restValues.alias || + mcpServer.alias || + "unknown"; + const payload: Record = { ...restValues, + ...stdioFields, + // Remove UI-only fields + stdio_config: undefined, + env_json: undefined, server_id: mcpServer.server_id, mcp_info: { - server_name: restValues.server_name || restValues.url, + server_name: mcpInfoServerName, description: restValues.description, mcp_server_cost_info: Object.keys(costConfig).length > 0 ? costConfig : null, }, @@ -386,7 +523,7 @@ const MCPServerEdit: React.FC = ({ }, ]} > - + = ({ }, ]} > - setAliasManuallyEdited(true)} /> + setAliasManuallyEdited(true)} + className="rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500" + /> - - - validateMCPServerUrl(value) }, - ]} - > - + - Server-Sent Events (SSE) HTTP - - - - - {shouldShowAuthValueField && ( + {/* URL/Auth fields are only applicable for HTTP/SSE */} + {!isStdioTransport && ( + validateMCPServerUrl(value) }, + ]} + > + + + )} + + {!isStdioTransport && ( + + + + )} + + {isStdioTransport && ( +
+

+ Configure the stdio transport used to launch the MCP server process. You can either fill in the fields + below or paste a JSON configuration. +

+ + + + + + + + + + Authorization URL Override (optional) + + + + + } + name="authorization_url" + > + + + + Token URL Override (optional) + + + + + } + name="token_url" + > + + + + Registration URL Override (optional) + + + + + } + name="registration_url" + > + + +
+

Use OAuth to fetch a fresh access token and temporarily save it in the session as the authentication value.

+ + {oauthError &&

{oauthError}

} + {oauthStatus === "success" && oauthTokenResponse?.access_token && ( +

+ Token fetched. Expires in {oauthTokenResponse.expires_in ?? "?"} seconds. +

+ )} +
+ )} {/* Permission Management / Access Control Section */} diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx index 960bd4a518..635c787f30 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx @@ -43,9 +43,11 @@ export const MCPServerView: React.FC = ({ onBack(); }; - const { maskedUrl, hasToken } = getMaskedAndFullUrl(mcpServer.url); + const urlValue = mcpServer.url ?? ""; + const { maskedUrl, hasToken } = urlValue ? getMaskedAndFullUrl(urlValue) : { maskedUrl: "—", hasToken: false }; - const renderUrlWithToggle = (url: string, showFull: boolean) => { + const renderUrlWithToggle = (url: string | null | undefined, showFull: boolean) => { + if (!url) return "—"; if (!hasToken) return url; return showFull ? url : maskedUrl; }; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 5cb840ec7d..ecc4171a8a 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -21,6 +21,7 @@ export const OAUTH_FLOW = { export const TRANSPORT = { SSE: "sse", HTTP: "http", + STDIO: "stdio", }; export const handleTransport = (transport?: string | null): string => { @@ -137,7 +138,11 @@ export interface MCPServer { server_name?: string | null; alias?: string | null; description?: string | null; - url: string; + /** + * Only required for HTTP/SSE transports. + * For `stdio`, the backend can return null/undefined. + */ + url?: string | null; transport?: string | null; auth_type?: string | null; authorization_url?: string | null; @@ -158,6 +163,11 @@ export interface MCPServer { allowed_tools?: string[]; allow_all_keys?: boolean; available_on_public_internet?: boolean; + + /** Stdio-only fields (present when transport === 'stdio') */ + command?: string | null; + args?: string[] | null; + env?: Record | null; } export interface MCPServerProps {