| 1 | package plugin |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "fmt" |
| 6 | "net/http" |
| 7 | "strings" |
| 8 | "sync" |
| 9 | "time" |
| 10 | |
| 11 | "golang.org/x/oauth2" |
| 12 | ) |
| 13 | |
| 14 | type mcpOAuthSDKRuntime struct { |
| 15 | mu sync.Mutex |
| 16 | fatalErr error |
| 17 | fatalErrReturns int |
| 18 | } |
| 19 | |
| 20 | func (c *mcpOAuthClient) oauthToken(ctx context.Context, forceRefresh bool) (*oauth2.Token, error) { |
| 21 | return c.oauthTokenAfterRejection(ctx, forceRefresh, "") |
| 22 | } |
| 23 | |
| 24 | func (c *mcpOAuthClient) oauthTokenAfterRejection(ctx context.Context, forceRefresh bool, rejectedAccessToken string) (*oauth2.Token, error) { |
| 25 | if c == nil { |
| 26 | return nil, nil |
| 27 | } |
| 28 | c.runtime.mu.Lock() |
| 29 | defer c.runtime.mu.Unlock() |
| 30 | if c.runtime.fatalErr != nil { |
| 31 | err := c.runtime.fatalErr |
| 32 | c.runtime.fatalErrReturns-- |
| 33 | if c.runtime.fatalErrReturns <= 0 { |
| 34 | c.runtime.fatalErr = nil |
| 35 | } |
| 36 | return nil, err |
| 37 | } |
| 38 | if forceRefresh && rejectedAccessToken == "" { |
| 39 | rejectedAccessToken = c.state.AccessToken |
| 40 | } |
| 41 | if forceRefresh && rejectedAccessToken != "" && c.state.AccessToken != rejectedAccessToken && oauthAccessTokenUsable(c.state, time.Now()) { |
| 42 | forceRefresh = false |
| 43 | } |
| 44 | needsRefresh := forceRefresh || (strings.TrimSpace(c.state.RefreshToken) != "" && !c.state.Expiry.IsZero() && time.Now().Add(30*time.Second).After(c.state.Expiry)) |
| 45 | if needsRefresh { |
| 46 | if err := c.refresh(ctx, forceRefresh, rejectedAccessToken); err != nil { |
| 47 | c.runtime.fatalErr = err |
| 48 | c.runtime.fatalErrReturns = 1 |
| 49 | return nil, err |
| 50 | } |
| 51 | } |
| 52 | if strings.TrimSpace(c.state.AccessToken) == "" { |
| 53 | return nil, nil |
| 54 | } |
| 55 | tokenType := strings.TrimSpace(c.state.TokenType) |
| 56 | if tokenType == "" { |
| 57 | tokenType = "Bearer" |
| 58 | } |
| 59 | if !strings.EqualFold(tokenType, "Bearer") { |
| 60 | return nil, fmt.Errorf("MCP OAuth: unsupported token type %q", tokenType) |
| 61 | } |
| 62 | return &oauth2.Token{ |
| 63 | AccessToken: c.state.AccessToken, |
| 64 | TokenType: tokenType, |
| 65 | RefreshToken: c.state.RefreshToken, |
| 66 | Expiry: c.state.Expiry, |
| 67 | }, nil |
| 68 | } |
| 69 | |
| 70 | func (c *mcpOAuthClient) canRefresh() bool { |
| 71 | return c != nil && strings.TrimSpace(c.state.RefreshToken) != "" && strings.TrimSpace(c.state.TokenEndpoint) != "" |
| 72 | } |
| 73 | |
| 74 | type mcpOAuthTokenSource struct { |
| 75 | ctx context.Context |
| 76 | client *mcpOAuthClient |
| 77 | } |
| 78 | |
| 79 | func (s *mcpOAuthTokenSource) Token() (*oauth2.Token, error) { |
| 80 | return s.client.oauthToken(s.ctx, false) |
| 81 | } |
| 82 | |
| 83 | // TokenSource implements auth.OAuthHandler for the official MCP Go SDK. |
| 84 | func (c *mcpOAuthClient) TokenSource(ctx context.Context) (oauth2.TokenSource, error) { |
| 85 | if c == nil { |
| 86 | return nil, nil |
| 87 | } |
| 88 | return &mcpOAuthTokenSource{ctx: ctx, client: c}, nil |
| 89 | } |
| 90 | |
| 91 | // Authorize handles the SDK's single retry after a 401/403 without starting an |
| 92 | // interactive browser flow from a background tool call. |
| 93 | func (c *mcpOAuthClient) Authorize(ctx context.Context, request *http.Request, response *http.Response) error { |
| 94 | if response != nil && response.Body != nil { |
| 95 | _ = response.Body.Close() |
| 96 | } |
| 97 | if c == nil { |
| 98 | return fmt.Errorf("MCP OAuth authorization is required") |
| 99 | } |
| 100 | c.runtime.mu.Lock() |
| 101 | canRefresh := c.canRefresh() |
| 102 | c.runtime.mu.Unlock() |
| 103 | if !canRefresh { |
| 104 | return fmt.Errorf("MCP OAuth authorization is required; authorize this MCP server") |
| 105 | } |
| 106 | rejectedAccessToken := "" |
| 107 | if request != nil { |
| 108 | scheme, token, ok := strings.Cut(strings.TrimSpace(request.Header.Get("Authorization")), " ") |
| 109 | if ok && strings.EqualFold(scheme, "Bearer") { |
| 110 | rejectedAccessToken = strings.TrimSpace(token) |
| 111 | } |
| 112 | } |
| 113 | _, err := c.oauthTokenAfterRejection(ctx, true, rejectedAccessToken) |
| 114 | return err |
| 115 | } |
| 116 |