224 lines
6.1 KiB
Go
224 lines
6.1 KiB
Go
package auth
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net"
|
|
"net/http"
|
|
"os"
|
|
"os/exec"
|
|
"time"
|
|
|
|
mcpauth "github.com/modelcontextprotocol/go-sdk/auth"
|
|
mcp "github.com/modelcontextprotocol/go-sdk/mcp"
|
|
"github.com/modelcontextprotocol/go-sdk/oauthex"
|
|
"golang.org/x/oauth2"
|
|
|
|
"s1d3sw1ped/robinhood-agentic-mcp/client"
|
|
"s1d3sw1ped/robinhood-agentic-mcp/internal/wire"
|
|
)
|
|
|
|
func oauthIdentity(cfg Config) (name, version string) {
|
|
cfg = cfg.WithDefaults()
|
|
return cfg.Name, cfg.Version
|
|
}
|
|
|
|
func loginOAuth(ctx context.Context, cfg Config) (string, error) {
|
|
cfg = cfg.WithDefaults()
|
|
name, version := oauthIdentity(cfg.WithDefaults())
|
|
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
return "", fmt.Errorf("listen: %w", err)
|
|
}
|
|
defer ln.Close()
|
|
port := ln.Addr().(*net.TCPAddr).Port
|
|
redirect := fmt.Sprintf("http://127.0.0.1:%d/callback", port)
|
|
|
|
codeCh := make(chan mcpauth.AuthorizationResult, 1)
|
|
mux := http.NewServeMux()
|
|
mux.HandleFunc("/callback", func(w http.ResponseWriter, r *http.Request) {
|
|
q := r.URL.Query()
|
|
if q.Get("error") != "" {
|
|
http.Error(w, q.Get("error_description"), http.StatusBadRequest)
|
|
select {
|
|
case codeCh <- mcpauth.AuthorizationResult{}:
|
|
default:
|
|
}
|
|
return
|
|
}
|
|
res := mcpauth.AuthorizationResult{Code: q.Get("code"), State: q.Get("state"), Iss: q.Get("iss")}
|
|
_, _ = w.Write([]byte(name + " is signed in. You can close this tab."))
|
|
select {
|
|
case codeCh <- res:
|
|
default:
|
|
}
|
|
})
|
|
srv := &http.Server{Handler: mux}
|
|
go func() { _ = srv.Serve(ln) }()
|
|
defer func() {
|
|
shctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
|
defer cancel()
|
|
_ = srv.Shutdown(shctx)
|
|
}()
|
|
|
|
var saved TokenSet
|
|
oauthCfg := &mcpauth.AuthorizationCodeHandlerConfig{
|
|
RedirectURL: redirect,
|
|
DynamicClientRegistrationConfig: &mcpauth.DynamicClientRegistrationConfig{
|
|
Metadata: &oauthex.ClientRegistrationMetadata{
|
|
RedirectURIs: []string{redirect},
|
|
ClientName: name,
|
|
ApplicationType: "native",
|
|
TokenEndpointAuthMethod: "none",
|
|
GrantTypes: []string{"authorization_code", "refresh_token"},
|
|
ResponseTypes: []string{"code"},
|
|
},
|
|
},
|
|
RequestRefreshToken: true,
|
|
AuthorizationCodeFetcher: func(ctx context.Context, args *mcpauth.AuthorizationArgs) (*mcpauth.AuthorizationResult, error) {
|
|
fmt.Fprintf(os.Stderr, "Open this URL to authorize %s with Robinhood:\n\n %s\n\n", name, args.URL)
|
|
_ = openBrowser(args.URL)
|
|
select {
|
|
case <-ctx.Done():
|
|
return nil, ctx.Err()
|
|
case res := <-codeCh:
|
|
if res.Code == "" {
|
|
return nil, fmt.Errorf("oauth callback missing code")
|
|
}
|
|
return &res, nil
|
|
}
|
|
},
|
|
NewTokenSource: func(ctx context.Context, oc *oauth2.Config, tok *oauth2.Token) (oauth2.TokenSource, error) {
|
|
saved = tokenSetFrom(oc, tok, redirect, saved.AccountID)
|
|
_ = WriteTokenSet(cfg.TokenFile, saved)
|
|
return oc.TokenSource(ctx, tok), nil
|
|
},
|
|
}
|
|
handler, err := mcpauth.NewAuthorizationCodeHandler(oauthCfg)
|
|
if err != nil {
|
|
return "", fmt.Errorf("oauth handler: %w", err)
|
|
}
|
|
transport := &mcp.StreamableClientTransport{
|
|
Endpoint: cfg.URL,
|
|
OAuthHandler: handler,
|
|
DisableStandaloneSSE: true,
|
|
}
|
|
mcpClient := mcp.NewClient(&mcp.Implementation{Name: name, Version: version}, nil)
|
|
ctx, cancel := context.WithTimeout(ctx, 5*time.Minute)
|
|
defer cancel()
|
|
sess, err := mcpClient.Connect(ctx, transport, nil)
|
|
if err != nil {
|
|
return "", fmt.Errorf("connect robinhood mcp: %w", err)
|
|
}
|
|
defer sess.Close()
|
|
|
|
c := &client.Client{URL: cfg.URL, Name: name, Version: version}
|
|
c.AttachSession(sess)
|
|
raw, err := c.Call(ctx, "get_accounts", map[string]any{})
|
|
if err != nil {
|
|
return "", fmt.Errorf("get_accounts: %w", err)
|
|
}
|
|
id, err := accountIDFrom(raw)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
saved.AccountID = id
|
|
if ts, err := handler.TokenSource(ctx); err == nil && ts != nil {
|
|
if tok, err := ts.Token(); err == nil {
|
|
saved.AccessToken = tok.AccessToken
|
|
saved.RefreshToken = tok.RefreshToken
|
|
saved.Expiry = tok.Expiry
|
|
saved.TokenType = tok.TokenType
|
|
}
|
|
}
|
|
if err := WriteTokenSet(cfg.TokenFile, saved); err != nil {
|
|
return "", err
|
|
}
|
|
return id, nil
|
|
}
|
|
|
|
func tokenSetFrom(oc *oauth2.Config, tok *oauth2.Token, redirect, accountID string) TokenSet {
|
|
if oc == nil {
|
|
oc = &oauth2.Config{}
|
|
}
|
|
if tok == nil {
|
|
tok = &oauth2.Token{}
|
|
}
|
|
return TokenSet{
|
|
AccessToken: tok.AccessToken,
|
|
RefreshToken: tok.RefreshToken,
|
|
TokenType: tok.TokenType,
|
|
Expiry: tok.Expiry,
|
|
ClientID: oc.ClientID,
|
|
ClientSecret: oc.ClientSecret,
|
|
AuthURL: oc.Endpoint.AuthURL,
|
|
TokenURL: oc.Endpoint.TokenURL,
|
|
RedirectURL: redirect,
|
|
AccountID: accountID,
|
|
}
|
|
}
|
|
|
|
func openBrowser(rawURL string) error {
|
|
cmds := [][]string{
|
|
{"xdg-open", rawURL},
|
|
{"gio", "open", rawURL},
|
|
{"open", rawURL},
|
|
}
|
|
for _, c := range cmds {
|
|
if err := exec.Command(c[0], c[1:]...).Start(); err == nil {
|
|
return nil
|
|
}
|
|
}
|
|
return fmt.Errorf("open browser")
|
|
}
|
|
|
|
func accountIDFrom(raw json.RawMessage) (string, error) {
|
|
data := wire.Unwrap(raw)
|
|
type row struct {
|
|
ID string `json:"id"`
|
|
AccountNumber string `json:"account_number"`
|
|
AgenticAllowed bool `json:"agentic_allowed"`
|
|
Agentic bool `json:"agentic"`
|
|
}
|
|
var wrap struct {
|
|
Accounts []row `json:"accounts"`
|
|
}
|
|
if err := json.Unmarshal(data, &wrap); err != nil || wrap.Accounts == nil {
|
|
var one row
|
|
if err2 := json.Unmarshal(data, &one); err2 != nil {
|
|
if err != nil {
|
|
return "", fmt.Errorf("parse accounts: %w", err)
|
|
}
|
|
return "", fmt.Errorf("parse accounts")
|
|
}
|
|
wrap.Accounts = []row{one}
|
|
}
|
|
for _, a := range wrap.Accounts {
|
|
id := a.AccountNumber
|
|
if id == "" {
|
|
id = a.ID
|
|
}
|
|
if a.AgenticAllowed || a.Agentic {
|
|
return id, nil
|
|
}
|
|
}
|
|
return "", fmt.Errorf("no agentic account")
|
|
}
|
|
|
|
// ClientToken maps a persisted TokenSet to the fields ConnectSession needs.
|
|
func ClientToken(t TokenSet) client.Token {
|
|
return client.Token{
|
|
AccessToken: t.AccessToken,
|
|
RefreshToken: t.RefreshToken,
|
|
TokenType: t.TokenType,
|
|
Expiry: t.Expiry,
|
|
ClientID: t.ClientID,
|
|
ClientSecret: t.ClientSecret,
|
|
AuthURL: t.AuthURL,
|
|
TokenURL: t.TokenURL,
|
|
RedirectURL: t.RedirectURL,
|
|
}
|
|
}
|