1// Package mcp provides functionality for managing Model Context Protocol (MCP)
2// clients within the Crush application.
3package mcp
4
5import (
6 "cmp"
7 "context"
8 "errors"
9 "fmt"
10 "io"
11 "log/slog"
12 "net/http"
13 "os"
14 "os/exec"
15 "strings"
16 "sync"
17 "time"
18
19 "github.com/charmbracelet/crush/internal/config"
20 "github.com/charmbracelet/crush/internal/csync"
21 "github.com/charmbracelet/crush/internal/home"
22 "github.com/charmbracelet/crush/internal/permission"
23 "github.com/charmbracelet/crush/internal/pubsub"
24 "github.com/charmbracelet/crush/internal/version"
25 "github.com/modelcontextprotocol/go-sdk/mcp"
26)
27
28func parseLevel(level mcp.LoggingLevel) slog.Level {
29 switch level {
30 case "info":
31 return slog.LevelInfo
32 case "notice":
33 return slog.LevelInfo
34 case "warning":
35 return slog.LevelWarn
36 default:
37 return slog.LevelDebug
38 }
39}
40
41var (
42 sessions = csync.NewMap[string, *mcp.ClientSession]()
43 states = csync.NewMap[string, ClientInfo]()
44 broker = pubsub.NewBroker[Event]()
45 initOnce sync.Once
46 initDone = make(chan struct{})
47)
48
49// State represents the current state of an MCP client
50type State int
51
52const (
53 StateDisabled State = iota
54 StateStarting
55 StateConnected
56 StateError
57)
58
59func (s State) String() string {
60 switch s {
61 case StateDisabled:
62 return "disabled"
63 case StateStarting:
64 return "starting"
65 case StateConnected:
66 return "connected"
67 case StateError:
68 return "error"
69 default:
70 return "unknown"
71 }
72}
73
74// EventType represents the type of MCP event
75type EventType uint
76
77const (
78 EventStateChanged EventType = iota
79 EventToolsListChanged
80 EventPromptsListChanged
81 EventResourcesListChanged
82)
83
84// Event represents an event in the MCP system
85type Event struct {
86 Type EventType
87 Name string
88 State State
89 Error error
90 Counts Counts
91}
92
93// Counts number of available tools, prompts, etc.
94type Counts struct {
95 Tools int
96 Prompts int
97 Resources int
98}
99
100// ClientInfo holds information about an MCP client's state
101type ClientInfo struct {
102 Name string
103 State State
104 Error error
105 Client *mcp.ClientSession
106 Counts Counts
107 ConnectedAt time.Time
108}
109
110// SubscribeEvents returns a channel for MCP events
111func SubscribeEvents(ctx context.Context) <-chan pubsub.Event[Event] {
112 return broker.Subscribe(ctx)
113}
114
115// GetStates returns the current state of all MCP clients
116func GetStates() map[string]ClientInfo {
117 return states.Copy()
118}
119
120// GetState returns the state of a specific MCP client
121func GetState(name string) (ClientInfo, bool) {
122 return states.Get(name)
123}
124
125// Close closes all MCP clients. This should be called during application shutdown.
126func Close() error {
127 ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
128 defer cancel()
129
130 var wg sync.WaitGroup
131 for name, session := range sessions.Seq2() {
132 wg.Go(func() {
133 done := make(chan error, 1)
134 go func() {
135 done <- session.Close()
136 }()
137 select {
138 case err := <-done:
139 if err != nil &&
140 !errors.Is(err, io.EOF) &&
141 !errors.Is(err, context.Canceled) &&
142 err.Error() != "signal: killed" {
143 slog.Warn("Failed to shutdown MCP client", "name", name, "error", err)
144 }
145 case <-ctx.Done():
146 }
147 })
148 }
149 wg.Wait()
150 broker.Shutdown()
151 return nil
152}
153
154// Initialize initializes MCP clients based on the provided configuration.
155func Initialize(ctx context.Context, permissions permission.Service, cfg *config.Config) {
156 slog.Info("Initializing MCP clients")
157 var wg sync.WaitGroup
158 // Initialize states for all configured MCPs
159 for name, m := range cfg.MCP {
160 if m.Disabled {
161 updateState(name, StateDisabled, nil, nil, Counts{})
162 slog.Debug("Skipping disabled MCP", "name", name)
163 continue
164 }
165
166 // Set initial starting state
167 updateState(name, StateStarting, nil, nil, Counts{})
168
169 wg.Add(1)
170 go func(name string, m config.MCPConfig) {
171 defer func() {
172 wg.Done()
173 if r := recover(); r != nil {
174 var err error
175 switch v := r.(type) {
176 case error:
177 err = v
178 case string:
179 err = fmt.Errorf("panic: %s", v)
180 default:
181 err = fmt.Errorf("panic: %v", v)
182 }
183 updateState(name, StateError, err, nil, Counts{})
184 slog.Error("Panic in MCP client initialization", "error", err, "name", name)
185 }
186 }()
187
188 // createSession handles its own timeout internally.
189 session, err := createSession(ctx, name, m, cfg.Resolver())
190 if err != nil {
191 return
192 }
193
194 tools, err := getTools(ctx, session)
195 if err != nil {
196 slog.Error("Error listing tools", "error", err)
197 updateState(name, StateError, err, nil, Counts{})
198 session.Close()
199 return
200 }
201
202 prompts, err := getPrompts(ctx, session)
203 if err != nil {
204 slog.Error("Error listing prompts", "error", err)
205 updateState(name, StateError, err, nil, Counts{})
206 session.Close()
207 return
208 }
209
210 resources, err := getResources(ctx, session)
211 if err != nil {
212 slog.Error("Error listing resources", "error", err)
213 updateState(name, StateError, err, nil, Counts{})
214 session.Close()
215 return
216 }
217
218 toolCount := updateTools(cfg, name, tools)
219 updatePrompts(name, prompts)
220 resourceCount := updateResources(name, resources)
221 sessions.Set(name, session)
222
223 updateState(name, StateConnected, nil, session, Counts{
224 Tools: toolCount,
225 Prompts: len(prompts),
226 Resources: resourceCount,
227 })
228 }(name, m)
229 }
230 wg.Wait()
231 initOnce.Do(func() { close(initDone) })
232}
233
234// WaitForInit blocks until MCP initialization is complete.
235// If Initialize was never called, this returns immediately.
236func WaitForInit(ctx context.Context) error {
237 select {
238 case <-initDone:
239 return nil
240 case <-ctx.Done():
241 return ctx.Err()
242 }
243}
244
245func getOrRenewClient(ctx context.Context, cfg *config.Config, name string) (*mcp.ClientSession, error) {
246 sess, ok := sessions.Get(name)
247 if !ok {
248 return nil, fmt.Errorf("mcp '%s' not available", name)
249 }
250
251 m := cfg.MCP[name]
252 state, _ := states.Get(name)
253
254 timeout := mcpTimeout(m)
255 pingCtx, cancel := context.WithTimeout(ctx, timeout)
256 defer cancel()
257 err := sess.Ping(pingCtx, nil)
258 if err == nil {
259 return sess, nil
260 }
261 updateState(name, StateError, maybeTimeoutErr(err, timeout), nil, state.Counts)
262
263 sess, err = createSession(ctx, name, m, cfg.Resolver())
264 if err != nil {
265 return nil, err
266 }
267
268 updateState(name, StateConnected, nil, sess, state.Counts)
269 sessions.Set(name, sess)
270 return sess, nil
271}
272
273// updateState updates the state of an MCP client and publishes an event
274func updateState(name string, state State, err error, client *mcp.ClientSession, counts Counts) {
275 info := ClientInfo{
276 Name: name,
277 State: state,
278 Error: err,
279 Client: client,
280 Counts: counts,
281 }
282 switch state {
283 case StateConnected:
284 info.ConnectedAt = time.Now()
285 case StateError:
286 sessions.Del(name)
287 }
288 states.Set(name, info)
289
290 // Publish state change event
291 broker.Publish(pubsub.UpdatedEvent, Event{
292 Type: EventStateChanged,
293 Name: name,
294 State: state,
295 Error: err,
296 Counts: counts,
297 })
298}
299
300func createSession(ctx context.Context, name string, m config.MCPConfig, resolver config.VariableResolver) (*mcp.ClientSession, error) {
301 timeout := mcpTimeout(m)
302 mcpCtx, cancel := context.WithCancel(ctx)
303 cancelTimer := time.AfterFunc(timeout, cancel)
304
305 transport, err := createTransport(mcpCtx, m, resolver)
306 if err != nil {
307 updateState(name, StateError, err, nil, Counts{})
308 slog.Error("Error creating MCP client", "error", err, "name", name)
309 cancel()
310 cancelTimer.Stop()
311 return nil, err
312 }
313
314 client := mcp.NewClient(
315 &mcp.Implementation{
316 Name: "crush",
317 Version: version.Version,
318 Title: "Crush",
319 },
320 &mcp.ClientOptions{
321 ToolListChangedHandler: func(context.Context, *mcp.ToolListChangedRequest) {
322 broker.Publish(pubsub.UpdatedEvent, Event{
323 Type: EventToolsListChanged,
324 Name: name,
325 })
326 },
327 PromptListChangedHandler: func(context.Context, *mcp.PromptListChangedRequest) {
328 broker.Publish(pubsub.UpdatedEvent, Event{
329 Type: EventPromptsListChanged,
330 Name: name,
331 })
332 },
333 ResourceListChangedHandler: func(context.Context, *mcp.ResourceListChangedRequest) {
334 broker.Publish(pubsub.UpdatedEvent, Event{
335 Type: EventResourcesListChanged,
336 Name: name,
337 })
338 },
339 LoggingMessageHandler: func(ctx context.Context, req *mcp.LoggingMessageRequest) {
340 level := parseLevel(req.Params.Level)
341 slog.Log(ctx, level, "MCP log", "name", name, "logger", req.Params.Logger, "data", req.Params.Data)
342 },
343 },
344 )
345
346 session, err := client.Connect(mcpCtx, transport, nil)
347 if err != nil {
348 err = maybeStdioErr(err, transport)
349 updateState(name, StateError, maybeTimeoutErr(err, timeout), nil, Counts{})
350 slog.Error("MCP client failed to initialize", "error", err, "name", name)
351 cancel()
352 cancelTimer.Stop()
353 return nil, err
354 }
355
356 cancelTimer.Stop()
357 slog.Debug("MCP client initialized", "name", name)
358 return session, nil
359}
360
361// maybeStdioErr if a stdio mcp prints an error in non-json format, it'll fail
362// to parse, and the cli will then close it, causing the EOF error.
363// so, if we got an EOF err, and the transport is STDIO, we try to exec it
364// again with a timeout and collect the output so we can add details to the
365// error.
366// this happens particularly when starting things with npx, e.g. if node can't
367// be found or some other error like that.
368func maybeStdioErr(err error, transport mcp.Transport) error {
369 if !errors.Is(err, io.EOF) {
370 return err
371 }
372 ct, ok := transport.(*mcp.CommandTransport)
373 if !ok {
374 return err
375 }
376 if err2 := stdioCheck(ct.Command); err2 != nil {
377 err = errors.Join(err, err2)
378 }
379 return err
380}
381
382func maybeTimeoutErr(err error, timeout time.Duration) error {
383 if errors.Is(err, context.Canceled) {
384 return fmt.Errorf("timed out after %s", timeout)
385 }
386 return err
387}
388
389func createTransport(ctx context.Context, m config.MCPConfig, resolver config.VariableResolver) (mcp.Transport, error) {
390 switch m.Type {
391 case config.MCPStdio:
392 command, err := resolver.ResolveValue(m.Command)
393 if err != nil {
394 return nil, fmt.Errorf("invalid mcp command: %w", err)
395 }
396 if strings.TrimSpace(command) == "" {
397 return nil, fmt.Errorf("mcp stdio config requires a non-empty 'command' field")
398 }
399 cmd := exec.CommandContext(ctx, home.Long(command), m.Args...)
400 cmd.Env = append(os.Environ(), m.ResolvedEnv()...)
401 return &mcp.CommandTransport{
402 Command: cmd,
403 }, nil
404 case config.MCPHttp:
405 if strings.TrimSpace(m.URL) == "" {
406 return nil, fmt.Errorf("mcp http config requires a non-empty 'url' field")
407 }
408 client := &http.Client{
409 Transport: &headerRoundTripper{
410 headers: m.ResolvedHeaders(),
411 },
412 }
413 return &mcp.StreamableClientTransport{
414 Endpoint: m.URL,
415 HTTPClient: client,
416 }, nil
417 case config.MCPSSE:
418 if strings.TrimSpace(m.URL) == "" {
419 return nil, fmt.Errorf("mcp sse config requires a non-empty 'url' field")
420 }
421 client := &http.Client{
422 Transport: &headerRoundTripper{
423 headers: m.ResolvedHeaders(),
424 },
425 }
426 return &mcp.SSEClientTransport{
427 Endpoint: m.URL,
428 HTTPClient: client,
429 }, nil
430 default:
431 return nil, fmt.Errorf("unsupported mcp type: %s", m.Type)
432 }
433}
434
435type headerRoundTripper struct {
436 headers map[string]string
437}
438
439func (rt headerRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
440 for k, v := range rt.headers {
441 req.Header.Set(k, v)
442 }
443 return http.DefaultTransport.RoundTrip(req)
444}
445
446func mcpTimeout(m config.MCPConfig) time.Duration {
447 return time.Duration(cmp.Or(m.Timeout, 15)) * time.Second
448}
449
450func stdioCheck(old *exec.Cmd) error {
451 ctx, cancel := context.WithTimeout(context.Background(), time.Second*5)
452 defer cancel()
453 cmd := exec.CommandContext(ctx, old.Path, old.Args...)
454 cmd.Env = old.Env
455 out, err := cmd.CombinedOutput()
456 if err == nil || errors.Is(ctx.Err(), context.DeadlineExceeded) {
457 return nil
458 }
459 return fmt.Errorf("%w: %s", err, string(out))
460}