init.go

  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}