package mcp import ( "context" "encoding/json" "fmt" "sync" ) // Client is an MCP client bound to a single upstream server via a Transport. // It performs the initialize handshake and exposes tools/list and tools/call. type Client struct { transport Transport info Implementation mu sync.RWMutex initialized bool serverInfo Implementation serverCaps Capabilities } // NewClient wraps a transport. clientInfo identifies Nexus to the upstream. func NewClient(t Transport, clientInfo Implementation) *Client { return &Client{transport: t, info: clientInfo} } // SetNotificationHandler forwards peer notifications (e.g. list_changed). func (c *Client) SetNotificationHandler(h func(*Message)) { c.transport.SetNotificationHandler(h) } // Initialize performs the MCP initialize handshake and sends the initialized // notification. It is safe to call once per client. func (c *Client) Initialize(ctx context.Context) error { params := InitializeParams{ ProtocolVersion: ProtocolVersion, Capabilities: Capabilities{}, ClientInfo: c.info, } resp, err := c.transport.Call(ctx, MethodInitialize, params) if err != nil { return fmt.Errorf("initialize: %w", err) } if resp.Error != nil { return fmt.Errorf("initialize: %w", resp.Error) } var res InitializeResult if err := resp.UnmarshalResult(&res); err != nil { return fmt.Errorf("initialize: decode result: %w", err) } c.mu.Lock() c.initialized = true c.serverInfo = res.ServerInfo c.serverCaps = res.Capabilities c.mu.Unlock() if err := c.transport.Notify(ctx, MethodInitialized, nil); err != nil { return fmt.Errorf("initialized notification: %w", err) } return nil } // ServerInfo returns the upstream's advertised identity (valid after Initialize). func (c *Client) ServerInfo() Implementation { c.mu.RLock() defer c.mu.RUnlock() return c.serverInfo } // ListTools returns all tools exposed by the upstream, following pagination. func (c *Client) ListTools(ctx context.Context) ([]ToolDefinition, error) { var all []ToolDefinition cursor := "" for { resp, err := c.transport.Call(ctx, MethodToolsList, ListToolsParams{Cursor: cursor}) if err != nil { return nil, fmt.Errorf("tools/list: %w", err) } if resp.Error != nil { return nil, fmt.Errorf("tools/list: %w", resp.Error) } var res ListToolsResult if err := resp.UnmarshalResult(&res); err != nil { return nil, fmt.Errorf("tools/list: decode: %w", err) } all = append(all, res.Tools...) if res.NextCursor == "" { break } cursor = res.NextCursor } return all, nil } // CallTool invokes a tool by its upstream (un-namespaced) name. The raw result // message is returned so the caller can pass MCP content through unmodified; // callToolResult is a decoded convenience view. func (c *Client) CallTool(ctx context.Context, name string, arguments json.RawMessage) (json.RawMessage, *RPCError, error) { resp, err := c.transport.Call(ctx, MethodToolsCall, CallToolParams{Name: name, Arguments: arguments}) if err != nil { return nil, nil, fmt.Errorf("tools/call %s: %w", name, err) } if resp.Error != nil { return nil, resp.Error, nil } return resp.Result, nil, nil } // Ping issues an MCP ping, used for liveness checks. func (c *Client) Ping(ctx context.Context) error { resp, err := c.transport.Call(ctx, MethodPing, struct{}{}) if err != nil { return err } if resp.Error != nil { return resp.Error } return nil } // Close shuts down the underlying transport. func (c *Client) Close() error { return c.transport.Close() }