| 1 | package plugin |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "fmt" |
| 6 | "net/url" |
| 7 | "strings" |
| 8 | "time" |
| 9 | ) |
| 10 | |
| 11 | func (c *mcpOAuthClient) refresh(ctx context.Context, force bool, rejectedAccessToken string) error { |
| 12 | releaseGate, err := acquireMCPOAuthRefreshGate(ctx, c.stateDir) |
| 13 | if err != nil { |
| 14 | return fmt.Errorf("serialize MCP OAuth token refresh: %w", err) |
| 15 | } |
| 16 | defer releaseGate() |
| 17 | |
| 18 | // The file lock protects snapshots; network I/O happens after it is released. |
| 19 | release, err := acquireMCPOAuthStateLock(ctx, c.stateDir) |
| 20 | if err != nil { |
| 21 | return fmt.Errorf("lock MCP OAuth token refresh: %w", err) |
| 22 | } |
| 23 | latest, err := loadMCPOAuthState(c.stateDir) |
| 24 | if err != nil { |
| 25 | release() |
| 26 | return err |
| 27 | } |
| 28 | if strings.TrimSpace(latest.Resource) != "" && !sameCanonicalResource(latest.Resource, c.state.Resource) { |
| 29 | release() |
| 30 | return fmt.Errorf("MCP OAuth token refresh: stored token belongs to a different MCP resource") |
| 31 | } |
| 32 | c.state = latest |
| 33 | if oauthAccessTokenUsable(latest, time.Now()) && (!force || rejectedAccessToken != "" && latest.AccessToken != rejectedAccessToken) { |
| 34 | release() |
| 35 | return nil |
| 36 | } |
| 37 | if !c.canRefresh() { |
| 38 | release() |
| 39 | return fmt.Errorf("MCP OAuth access token expired and no refresh token is available; authorize again") |
| 40 | } |
| 41 | refreshState := latest |
| 42 | generation, err := loadMCPOAuthGeneration(c.stateDir) |
| 43 | if err != nil { |
| 44 | release() |
| 45 | return err |
| 46 | } |
| 47 | release() |
| 48 | |
| 49 | form := url.Values{ |
| 50 | "grant_type": {"refresh_token"}, |
| 51 | "refresh_token": {refreshState.RefreshToken}, |
| 52 | "client_id": {refreshState.ClientID}, |
| 53 | "resource": {refreshState.Resource}, |
| 54 | } |
| 55 | if refreshState.Scope != "" { |
| 56 | form.Set("scope", refreshState.Scope) |
| 57 | } |
| 58 | token, err := requestOAuthToken(ctx, c.client, refreshState, form) |
| 59 | if err != nil { |
| 60 | return fmt.Errorf("refresh MCP OAuth token: %w", err) |
| 61 | } |
| 62 | |
| 63 | release, err = acquireMCPOAuthStateLock(ctx, c.stateDir) |
| 64 | if err != nil { |
| 65 | return fmt.Errorf("lock MCP OAuth token refresh result: %w", err) |
| 66 | } |
| 67 | defer release() |
| 68 | currentGeneration, err := loadMCPOAuthGeneration(c.stateDir) |
| 69 | if err != nil { |
| 70 | return err |
| 71 | } |
| 72 | current, err := loadMCPOAuthState(c.stateDir) |
| 73 | if err != nil { |
| 74 | return err |
| 75 | } |
| 76 | if currentGeneration != generation || !sameOAuthRefreshState(current, refreshState) { |
| 77 | if currentGeneration != generation { |
| 78 | return fmt.Errorf("MCP OAuth token refresh was invalidated while contacting the token endpoint; authorize again") |
| 79 | } |
| 80 | if oauthAccessTokenUsable(current, time.Now()) { |
| 81 | c.state = current |
| 82 | return nil |
| 83 | } |
| 84 | return fmt.Errorf("MCP OAuth token state changed while refreshing; authorize again") |
| 85 | } |
| 86 | oldRefresh := refreshState.RefreshToken |
| 87 | applyTokenResponse(&refreshState, token, time.Now()) |
| 88 | if refreshState.RefreshToken == "" { |
| 89 | refreshState.RefreshToken = oldRefresh |
| 90 | } |
| 91 | if err := saveMCPOAuthState(c.stateDir, refreshState); err != nil { |
| 92 | return err |
| 93 | } |
| 94 | c.state = refreshState |
| 95 | return nil |
| 96 | } |
| 97 | |
| 98 | func sameOAuthRefreshState(a, b mcpOAuthState) bool { |
| 99 | return a.Version == b.Version && |
| 100 | a.Resource == b.Resource && |
| 101 | a.Issuer == b.Issuer && |
| 102 | a.AuthorizationEndpoint == b.AuthorizationEndpoint && |
| 103 | a.TokenEndpoint == b.TokenEndpoint && |
| 104 | a.RegistrationEndpoint == b.RegistrationEndpoint && |
| 105 | a.ClientID == b.ClientID && |
| 106 | a.ClientSecret == b.ClientSecret && |
| 107 | a.TokenEndpointAuthMethod == b.TokenEndpointAuthMethod && |
| 108 | a.Scope == b.Scope && |
| 109 | a.AccessToken == b.AccessToken && |
| 110 | a.RefreshToken == b.RefreshToken && |
| 111 | a.TokenType == b.TokenType && |
| 112 | a.Expiry.Equal(b.Expiry) |
| 113 | } |
| 114 | |
| 115 | func oauthAccessTokenUsable(state mcpOAuthState, now time.Time) bool { |
| 116 | return strings.TrimSpace(state.AccessToken) != "" && (state.Expiry.IsZero() || now.Add(30*time.Second).Before(state.Expiry)) |
| 117 | } |
| 118 |