Files
robinhood-agentic-mcp/client/session.go
T
ash bc50f3fd6b
CI / Test and build (pull_request) Successful in 12s
client: Prefer structuredContent in toolJSON
get_equity_historicals can return MCP StructuredContent with empty or
non-JSON TextContent. Prefer StructuredContent (marshal to RawMessage)
so paper tradey once/run can parse historicals.

#13
2026-09-05 04:40:48 +00:00

142 lines
3.9 KiB
Go

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
}
if res.StructuredContent != nil {
b, err := json.Marshal(res.StructuredContent)
if err != nil {
return nil, err
}
return json.RawMessage(b), nil
}
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)
}