feat: add OAuth login and MCP session connect
This commit is contained in:
+6
-1
@@ -4,6 +4,8 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
|
||||
mcp "github.com/modelcontextprotocol/go-sdk/mcp"
|
||||
)
|
||||
|
||||
// Client is a Robinhood Agentic MCP client.
|
||||
@@ -11,7 +13,7 @@ type Client struct {
|
||||
URL, Token, Name, Version string
|
||||
HTTP *http.Client
|
||||
Hook Caller
|
||||
session any
|
||||
session *mcp.ClientSession
|
||||
}
|
||||
|
||||
// Call invokes an MCP tool by name and returns its JSON result.
|
||||
@@ -19,5 +21,8 @@ func (c *Client) Call(ctx context.Context, name string, args map[string]any) (js
|
||||
if c.Hook != nil {
|
||||
return c.Hook.Call(ctx, name, args)
|
||||
}
|
||||
if c.session != nil {
|
||||
return c.sessionCall(ctx, name, args)
|
||||
}
|
||||
return c.rpcCall(ctx, name, args)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
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)
|
||||
}
|
||||
Reference in New Issue
Block a user