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 {