返回 DeepSeek-Reasonix
oauth_sdk.go
根目录 / internal / plugin / oauth_sdk.go
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
116 lines GO