package client import ( "context" "encoding/json" "fmt" "net/http" "time" "github.com/modelcontextprotocol/go-sdk/auth" mcp "github.com/modelcontextprotocol/go-sdk/mcp" "github.com/modelcontextprotocol/go-sdk/oauthex" "golang.org/x/oauth2" ) // Token holds the OAuth/bearer fields ConnectSession needs. type Token struct { AccessToken, RefreshToken, TokenType string Expiry time.Time ClientID, ClientSecret string AuthURL, TokenURL, RedirectURL string } // ConnectSession opens a streamable MCP session using tok. func ConnectSession(ctx context.Context, url string, tok Token, name, version string) (*Client, error) { sess, err := connectSession(ctx, url, tok, name, version) if err != nil { return nil, err } return &Client{URL: url, Token: tok.AccessToken, Name: name, Version: version, session: sess}, nil } // AttachSession sets the SDK session used by Call when Hook is nil. func (c *Client) AttachSession(sess *mcp.ClientSession) { c.session = sess } func connectSession(ctx context.Context, mcpURL string, tok Token, name, version string) (*mcp.ClientSession, error) { var handler auth.OAuthHandler if tok.ClientID != "" && tok.TokenURL != "" { oc := &oauth2.Config{ ClientID: tok.ClientID, ClientSecret: tok.ClientSecret, RedirectURL: tok.RedirectURL, Endpoint: oauth2.Endpoint{AuthURL: tok.AuthURL, TokenURL: tok.TokenURL}, } ot := &oauth2.Token{ AccessToken: tok.AccessToken, RefreshToken: tok.RefreshToken, TokenType: tok.TokenType, Expiry: tok.Expiry, } h, err := auth.NewAuthorizationCodeHandler(&auth.AuthorizationCodeHandlerConfig{ RedirectURL: tok.RedirectURL, PreregisteredClient: &oauthex.ClientCredentials{ ClientID: tok.ClientID, }, InitialTokenSource: oc.TokenSource(ctx, ot), AuthorizationCodeFetcher: func(context.Context, *auth.AuthorizationArgs) (*auth.AuthorizationResult, error) { return nil, fmt.Errorf("robinhood session expired; run Login again") }, }) if err != nil { return nil, err } handler = h } httpClient := &http.Client{Timeout: 60 * time.Second} if handler == nil && tok.AccessToken != "" { httpClient.Transport = bearerRT{token: tok.AccessToken, base: http.DefaultTransport} } t := &mcp.StreamableClientTransport{ Endpoint: mcpURL, HTTPClient: httpClient, OAuthHandler: handler, DisableStandaloneSSE: true, } cli := mcp.NewClient(&mcp.Implementation{Name: name, Version: version}, nil) return cli.Connect(ctx, t, nil) } func (c *Client) sessionCall(ctx context.Context, name string, args map[string]any) (json.RawMessage, error) { if c.session == nil { return nil, ToolErrorf(name, "session is nil") } res, err := c.session.CallTool(ctx, &mcp.CallToolParams{Name: name, Arguments: args}) if err != nil { return nil, ToolErrorf(name, "%w", err) } raw, err := toolJSON(res) if err != nil { return nil, ToolErrorf(name, "%w", err) } return raw, nil } func toolJSON(res *mcp.CallToolResult) (json.RawMessage, error) { if res == nil { return nil, fmt.Errorf("empty tool result") } if err := res.GetError(); err != nil { return nil, err } var b []byte for _, c := range res.Content { t, ok := c.(*mcp.TextContent) if !ok { continue } b = append(b, t.Text...) } if len(b) == 0 { return json.RawMessage(`{}`), nil } if !json.Valid(b) { return nil, fmt.Errorf("non-json tool result") } return json.RawMessage(b), nil } type bearerRT struct { token string base http.RoundTripper } func (b bearerRT) RoundTrip(req *http.Request) (*http.Response, error) { r := req.Clone(req.Context()) r.Header.Set("Authorization", "Bearer "+b.token) base := b.base if base == nil { base = http.DefaultTransport } return base.RoundTrip(r) }