fix(go.sum): update ResolveSpec dependency to v1.0.87
This commit is contained in:
+216
@@ -0,0 +1,216 @@
|
||||
The MCP project is undergoing a licensing transition from the MIT License to the Apache License, Version 2.0 ("Apache-2.0"). All new code and specification contributions to the project are licensed under Apache-2.0. Documentation contributions (excluding specifications) are licensed under CC-BY-4.0.
|
||||
|
||||
Contributions for which relicensing consent has been obtained are licensed under Apache-2.0. Contributions made by authors who originally licensed their work under the MIT License and who have not yet granted explicit permission to relicense remain licensed under the MIT License.
|
||||
|
||||
No rights beyond those granted by the applicable original license are conveyed for such contributions.
|
||||
|
||||
---
|
||||
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright
|
||||
owner or by an individual or Legal Entity authorized to submit on behalf
|
||||
of the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
---
|
||||
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2024-2025 Model Context Protocol a Series of LF Projects, LLC.
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
|
||||
---
|
||||
|
||||
Creative Commons Attribution 4.0 International (CC-BY-4.0)
|
||||
|
||||
Documentation in this project (excluding specifications) is licensed under
|
||||
CC-BY-4.0. See https://creativecommons.org/licenses/by/4.0/legalcode for
|
||||
the full license text.
|
||||
+170
@@ -0,0 +1,170 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/oauthex"
|
||||
)
|
||||
|
||||
// TokenInfo holds information from a bearer token.
|
||||
type TokenInfo struct {
|
||||
Scopes []string
|
||||
Expiration time.Time
|
||||
// UserID is an optional identifier for the authenticated user.
|
||||
// If set by a TokenVerifier, it can be used by transports to prevent
|
||||
// session hijacking by ensuring that all requests for a given session
|
||||
// come from the same user.
|
||||
UserID string
|
||||
Extra map[string]any
|
||||
}
|
||||
|
||||
// The error that a TokenVerifier should return if the token cannot be verified.
|
||||
var ErrInvalidToken = errors.New("invalid token")
|
||||
|
||||
// The error that a TokenVerifier should return for OAuth-specific protocol errors.
|
||||
var ErrOAuth = errors.New("oauth error")
|
||||
|
||||
// A TokenVerifier checks the validity of a bearer token, and extracts information
|
||||
// from it. If verification fails, it should return an error that unwraps to ErrInvalidToken.
|
||||
// The HTTP request is provided in case verifying the token involves checking it.
|
||||
type TokenVerifier func(ctx context.Context, token string, req *http.Request) (*TokenInfo, error)
|
||||
|
||||
// RequireBearerTokenOptions are options for [RequireBearerToken].
|
||||
type RequireBearerTokenOptions struct {
|
||||
// The URL for the resource server metadata OAuth flow, to be returned as part
|
||||
// of the WWW-Authenticate header.
|
||||
ResourceMetadataURL string
|
||||
// The required scopes.
|
||||
Scopes []string
|
||||
}
|
||||
|
||||
type tokenInfoKey struct{}
|
||||
|
||||
// TokenInfoFromContext returns the [TokenInfo] stored in ctx, or nil if none.
|
||||
func TokenInfoFromContext(ctx context.Context) *TokenInfo {
|
||||
ti := ctx.Value(tokenInfoKey{})
|
||||
if ti == nil {
|
||||
return nil
|
||||
}
|
||||
return ti.(*TokenInfo)
|
||||
}
|
||||
|
||||
// RequireBearerToken returns a piece of middleware that verifies a bearer token using the verifier.
|
||||
// If verification succeeds, the [TokenInfo] is added to the request's context and the request proceeds.
|
||||
// If verification fails, the request fails with a 401 Unauthenticated, and the WWW-Authenticate header
|
||||
// is populated to enable [protected resource metadata].
|
||||
//
|
||||
// [protected resource metadata]: https://datatracker.ietf.org/doc/rfc9728
|
||||
func RequireBearerToken(verifier TokenVerifier, opts *RequireBearerTokenOptions) func(http.Handler) http.Handler {
|
||||
// Based on typescript-sdk/src/server/auth/middleware/bearerAuth.ts.
|
||||
|
||||
return func(handler http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
tokenInfo, errmsg, code := verify(r, verifier, opts)
|
||||
if code != 0 {
|
||||
if code == http.StatusUnauthorized || code == http.StatusForbidden {
|
||||
if opts != nil && opts.ResourceMetadataURL != "" {
|
||||
w.Header().Add("WWW-Authenticate", "Bearer resource_metadata="+opts.ResourceMetadataURL)
|
||||
}
|
||||
}
|
||||
http.Error(w, errmsg, code)
|
||||
return
|
||||
}
|
||||
r = r.WithContext(context.WithValue(r.Context(), tokenInfoKey{}, tokenInfo))
|
||||
handler.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func verify(req *http.Request, verifier TokenVerifier, opts *RequireBearerTokenOptions) (_ *TokenInfo, errmsg string, code int) {
|
||||
// Extract bearer token.
|
||||
authHeader := req.Header.Get("Authorization")
|
||||
fields := strings.Fields(authHeader)
|
||||
if len(fields) != 2 || strings.ToLower(fields[0]) != "bearer" {
|
||||
return nil, "no bearer token", http.StatusUnauthorized
|
||||
}
|
||||
|
||||
// Verify the token and get information from it.
|
||||
tokenInfo, err := verifier(req.Context(), fields[1], req)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrInvalidToken) {
|
||||
return nil, err.Error(), http.StatusUnauthorized
|
||||
}
|
||||
if errors.Is(err, ErrOAuth) {
|
||||
return nil, err.Error(), http.StatusBadRequest
|
||||
}
|
||||
return nil, err.Error(), http.StatusInternalServerError
|
||||
}
|
||||
if tokenInfo == nil {
|
||||
return nil, "token validation failed", http.StatusInternalServerError
|
||||
}
|
||||
|
||||
// Check scopes. All must be present.
|
||||
if opts != nil {
|
||||
// Note: quadratic, but N is small.
|
||||
for _, s := range opts.Scopes {
|
||||
if !slices.Contains(tokenInfo.Scopes, s) {
|
||||
return nil, "insufficient scope", http.StatusForbidden
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Check expiration.
|
||||
if tokenInfo.Expiration.IsZero() {
|
||||
return nil, "token missing expiration", http.StatusUnauthorized
|
||||
}
|
||||
if tokenInfo.Expiration.Before(time.Now()) {
|
||||
return nil, "token expired", http.StatusUnauthorized
|
||||
}
|
||||
return tokenInfo, "", 0
|
||||
}
|
||||
|
||||
// ProtectedResourceMetadataHandler returns an http.Handler that serves OAuth 2.0
|
||||
// protected resource metadata (RFC 9728) with CORS support.
|
||||
//
|
||||
// This handler allows cross-origin requests from any origin (Access-Control-Allow-Origin: *)
|
||||
// because OAuth metadata is public information intended for client discovery (RFC 9728 §3.1).
|
||||
// The metadata contains only non-sensitive configuration data about authorization servers
|
||||
// and supported scopes.
|
||||
//
|
||||
// No validation of metadata fields is performed; ensure metadata accuracy at configuration time.
|
||||
//
|
||||
// For more sophisticated CORS policies or to restrict origins, wrap this handler with a
|
||||
// CORS middleware like github.com/rs/cors or github.com/jub0bs/cors.
|
||||
func ProtectedResourceMetadataHandler(metadata *oauthex.ProtectedResourceMetadata) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
// Set CORS headers for cross-origin client discovery.
|
||||
// OAuth metadata is public information, so allowing any origin is safe.
|
||||
w.Header().Set("Access-Control-Allow-Origin", "*")
|
||||
w.Header().Set("Access-Control-Allow-Methods", "GET, OPTIONS")
|
||||
w.Header().Set("Access-Control-Allow-Headers", "Content-Type")
|
||||
|
||||
// Handle CORS preflight requests
|
||||
if r.Method == http.MethodOptions {
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
return
|
||||
}
|
||||
|
||||
// Only GET allowed for metadata retrieval
|
||||
if r.Method != http.MethodGet {
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if err := json.NewEncoder(w).Encode(metadata); err != nil {
|
||||
http.Error(w, "Failed to encode metadata", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
})
|
||||
}
|
||||
+565
@@ -0,0 +1,565 @@
|
||||
// Copyright 2026 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by the license
|
||||
// that can be found in the LICENSE file.
|
||||
|
||||
//go:build mcp_go_client_oauth
|
||||
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/oauthex"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
// ClientSecretAuthConfig is used to configure client authentication using client_secret.
|
||||
// Authentication method will be selected based on the authorization server's supported methods,
|
||||
// according to the following preference order:
|
||||
// 1. client_secret_post
|
||||
// 2. client_secret_basic
|
||||
type ClientSecretAuthConfig struct {
|
||||
// ClientID is the client ID to be used for client authentication.
|
||||
ClientID string
|
||||
// ClientSecret is the client secret to be used for client authentication.
|
||||
ClientSecret string
|
||||
}
|
||||
|
||||
// ClientIDMetadataDocumentConfig is used to configure the Client ID Metadata Document
|
||||
// based client registration per
|
||||
// https://modelcontextprotocol.io/specification/2025-11-25/basic/authorization#client-id-metadata-documents.
|
||||
// See https://client.dev/ for more information.
|
||||
type ClientIDMetadataDocumentConfig struct {
|
||||
// URL is the client identifier URL as per
|
||||
// https://datatracker.ietf.org/doc/html/draft-ietf-oauth-client-id-metadata-document-00#section-3.
|
||||
URL string
|
||||
}
|
||||
|
||||
// PreregisteredClientConfig is used to configure a pre-registered client per
|
||||
// https://modelcontextprotocol.io/specification/2025-11-25/basic/authorization#preregistration.
|
||||
// Currently only "client_secret_basic" and "client_secret_post" authentication methods are supported.
|
||||
type PreregisteredClientConfig struct {
|
||||
// ClientSecretAuthConfig is the client_secret based configuration to be used for client authentication.
|
||||
ClientSecretAuthConfig *ClientSecretAuthConfig
|
||||
}
|
||||
|
||||
// DynamicClientRegistrationConfig is used to configure dynamic client registration per
|
||||
// https://modelcontextprotocol.io/specification/2025-11-25/basic/authorization#dynamic-client-registration.
|
||||
type DynamicClientRegistrationConfig struct {
|
||||
// Metadata to be used in dynamic client registration request as per
|
||||
// https://datatracker.ietf.org/doc/html/rfc7591#section-2.
|
||||
Metadata *oauthex.ClientRegistrationMetadata
|
||||
}
|
||||
|
||||
// AuthorizationResult is the result of an authorization flow.
|
||||
// It is returned by [AuthorizationCodeHandler].AuthorizationCodeFetcher implementations.
|
||||
type AuthorizationResult struct {
|
||||
// Code is the authorization code obtained from the authorization server.
|
||||
Code string
|
||||
// State string returned by the authorization server.
|
||||
State string
|
||||
}
|
||||
|
||||
// AuthorizationArgs is the input to [AuthorizationCodeHandlerConfig].AuthorizationCodeFetcher.
|
||||
type AuthorizationArgs struct {
|
||||
// Authorization URL to be opened in a browser for the user to start the authorization process.
|
||||
URL string
|
||||
}
|
||||
|
||||
// AuthorizationCodeHandlerConfig is the configuration for [AuthorizationCodeHandler].
|
||||
type AuthorizationCodeHandlerConfig struct {
|
||||
// Client registration configuration.
|
||||
// It is attempted in the following order:
|
||||
// 1. Client ID Metadata Document
|
||||
// 2. Preregistration
|
||||
// 3. Dynamic Client Registration
|
||||
// At least one method must be configured.
|
||||
ClientIDMetadataDocumentConfig *ClientIDMetadataDocumentConfig
|
||||
PreregisteredClientConfig *PreregisteredClientConfig
|
||||
DynamicClientRegistrationConfig *DynamicClientRegistrationConfig
|
||||
|
||||
// RedirectURL is a required URL to redirect to after authorization.
|
||||
// The caller is responsible for handling the redirect out of band.
|
||||
//
|
||||
// If Dynamic Client Registration is used:
|
||||
// - this field is permitted to be empty, in which case it will be set
|
||||
// to the first redirect URI from
|
||||
// DynamicClientRegistrationConfig.Metadata.RedirectURIs.
|
||||
// - if the field is not empty, it must be one of the redirect URIs in
|
||||
// DynamicClientRegistrationConfig.Metadata.RedirectURIs.
|
||||
RedirectURL string
|
||||
|
||||
// AuthorizationCodeFetcher is a required function called to initiate the authorization flow.
|
||||
// It is responsible for opening the URL in a browser for the user to start the authorization process.
|
||||
// It should return the authorization code and state once the Authorization Server
|
||||
// redirects back to the RedirectURL.
|
||||
AuthorizationCodeFetcher func(ctx context.Context, args *AuthorizationArgs) (*AuthorizationResult, error)
|
||||
|
||||
// Client is an optional HTTP client to use for HTTP requests.
|
||||
// It is used for the following requests:
|
||||
// - Fetching Protected Resource Metadata
|
||||
// - Fetching Authorization Server Metadata
|
||||
// - Registering a client dynamically
|
||||
// - Exchanging an authorization code for an access token
|
||||
// - Refreshing an access token
|
||||
// Custom clients can include additional security configurations,
|
||||
// such as SSRF protections, see
|
||||
// https://modelcontextprotocol.io/docs/tutorials/security/security_best_practices#server-side-request-forgery-ssrf
|
||||
// If not provided, http.DefaultClient will be used.
|
||||
Client *http.Client
|
||||
}
|
||||
|
||||
// AuthorizationCodeHandler is an implementation of [OAuthHandler] that uses
|
||||
// the authorization code flow to obtain access tokens.
|
||||
type AuthorizationCodeHandler struct {
|
||||
config *AuthorizationCodeHandlerConfig
|
||||
|
||||
// tokenSource is the token source to use for authorization.
|
||||
tokenSource oauth2.TokenSource
|
||||
}
|
||||
|
||||
var _ OAuthHandler = (*AuthorizationCodeHandler)(nil)
|
||||
|
||||
func (h *AuthorizationCodeHandler) isOAuthHandler() {}
|
||||
|
||||
func (h *AuthorizationCodeHandler) TokenSource(ctx context.Context) (oauth2.TokenSource, error) {
|
||||
return h.tokenSource, nil
|
||||
}
|
||||
|
||||
// NewAuthorizationCodeHandler creates a new AuthorizationCodeHandler.
|
||||
// It performs validation of the configuration and returns an error if it is invalid.
|
||||
// The passed config is consumed by the handler and should not be modified after.
|
||||
func NewAuthorizationCodeHandler(config *AuthorizationCodeHandlerConfig) (*AuthorizationCodeHandler, error) {
|
||||
if config == nil {
|
||||
return nil, errors.New("config must be provided")
|
||||
}
|
||||
if config.ClientIDMetadataDocumentConfig == nil &&
|
||||
config.PreregisteredClientConfig == nil &&
|
||||
config.DynamicClientRegistrationConfig == nil {
|
||||
return nil, errors.New("at least one client registration configuration must be provided")
|
||||
}
|
||||
if config.AuthorizationCodeFetcher == nil {
|
||||
return nil, errors.New("AuthorizationCodeFetcher is required")
|
||||
}
|
||||
if config.ClientIDMetadataDocumentConfig != nil && !isNonRootHTTPSURL(config.ClientIDMetadataDocumentConfig.URL) {
|
||||
return nil, fmt.Errorf("client ID metadata document URL must be a non-root HTTPS URL")
|
||||
}
|
||||
preCfg := config.PreregisteredClientConfig
|
||||
if preCfg != nil {
|
||||
if preCfg.ClientSecretAuthConfig == nil {
|
||||
return nil, errors.New("ClientSecretAuthConfig is required for pre-registered client")
|
||||
}
|
||||
if preCfg.ClientSecretAuthConfig.ClientID == "" || preCfg.ClientSecretAuthConfig.ClientSecret == "" {
|
||||
return nil, fmt.Errorf("pre-registered client ID or secret is empty")
|
||||
}
|
||||
}
|
||||
dCfg := config.DynamicClientRegistrationConfig
|
||||
if dCfg != nil {
|
||||
if dCfg.Metadata == nil {
|
||||
return nil, errors.New("Metadata is required for dynamic client registration")
|
||||
}
|
||||
if len(dCfg.Metadata.RedirectURIs) == 0 {
|
||||
return nil, errors.New("Metadata.RedirectURIs is required for dynamic client registration")
|
||||
}
|
||||
if config.RedirectURL == "" {
|
||||
config.RedirectURL = dCfg.Metadata.RedirectURIs[0]
|
||||
} else if !slices.Contains(dCfg.Metadata.RedirectURIs, config.RedirectURL) {
|
||||
return nil, fmt.Errorf("RedirectURL %q is not in the list of allowed redirect URIs for dynamic client registration", config.RedirectURL)
|
||||
}
|
||||
}
|
||||
if config.RedirectURL == "" {
|
||||
// If the RedirectURL was supposed to be set by the dynamic client registration,
|
||||
// it should have been set by now. Otherwise, it is required.
|
||||
return nil, errors.New("RedirectURL is required")
|
||||
}
|
||||
if config.Client == nil {
|
||||
config.Client = http.DefaultClient
|
||||
}
|
||||
return &AuthorizationCodeHandler{config: config}, nil
|
||||
}
|
||||
|
||||
func isNonRootHTTPSURL(u string) bool {
|
||||
pu, err := url.Parse(u)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return pu.Scheme == "https" && pu.Path != ""
|
||||
}
|
||||
|
||||
// Authorize performs the authorization flow.
|
||||
// It is designed to perform the whole Authorization Code Grant flow.
|
||||
// On success, [AuthorizationCodeHandler.TokenSource] will return a token source with the fetched token.
|
||||
func (h *AuthorizationCodeHandler) Authorize(ctx context.Context, req *http.Request, resp *http.Response) error {
|
||||
defer resp.Body.Close()
|
||||
|
||||
wwwChallenges, err := oauthex.ParseWWWAuthenticate(resp.Header[http.CanonicalHeaderKey("WWW-Authenticate")])
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to parse WWW-Authenticate header: %v", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode == http.StatusForbidden && errorFromChallenges(wwwChallenges) != "insufficient_scope" {
|
||||
// We only want to perform step-up authorization for insufficient_scope errors.
|
||||
// Returning nil, so that the call is retried immediately and the response
|
||||
// is handled appropriately by the connection.
|
||||
// Step-up authorization is defined at
|
||||
// https://modelcontextprotocol.io/specification/2025-11-25/basic/authorization#step-up-authorization-flow
|
||||
return nil
|
||||
}
|
||||
|
||||
prm, err := h.getProtectedResourceMetadata(ctx, wwwChallenges, req.URL.String())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
asm, err := h.getAuthServerMetadata(ctx, prm)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
resolvedClientConfig, err := h.handleRegistration(ctx, asm)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
scps := scopesFromChallenges(wwwChallenges)
|
||||
if len(scps) == 0 && len(prm.ScopesSupported) > 0 {
|
||||
scps = prm.ScopesSupported
|
||||
}
|
||||
|
||||
cfg := &oauth2.Config{
|
||||
ClientID: resolvedClientConfig.clientID,
|
||||
ClientSecret: resolvedClientConfig.clientSecret,
|
||||
|
||||
Endpoint: oauth2.Endpoint{
|
||||
AuthURL: asm.AuthorizationEndpoint,
|
||||
TokenURL: asm.TokenEndpoint,
|
||||
AuthStyle: resolvedClientConfig.authStyle,
|
||||
},
|
||||
RedirectURL: h.config.RedirectURL,
|
||||
Scopes: scps,
|
||||
}
|
||||
|
||||
authRes, err := h.getAuthorizationCode(ctx, cfg, req.URL.String())
|
||||
if err != nil {
|
||||
// Purposefully leaving the error unwrappable so it can be handled by the caller.
|
||||
return err
|
||||
}
|
||||
|
||||
return h.exchangeAuthorizationCode(ctx, cfg, authRes, prm.Resource)
|
||||
}
|
||||
|
||||
// resourceMetadataURLFromChallenges returns a resource metadata URL from the given "WWW-Authenticate" header challenges,
|
||||
// or the empty string if there is none.
|
||||
func resourceMetadataURLFromChallenges(cs []oauthex.Challenge) string {
|
||||
for _, c := range cs {
|
||||
if u := c.Params["resource_metadata"]; u != "" {
|
||||
return u
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// scopesFromChallenges returns the scopes from the given "WWW-Authenticate" header challenges.
|
||||
// It only looks at challenges with the "Bearer" scheme.
|
||||
func scopesFromChallenges(cs []oauthex.Challenge) []string {
|
||||
for _, c := range cs {
|
||||
if c.Scheme == "bearer" && c.Params["scope"] != "" {
|
||||
return strings.Fields(c.Params["scope"])
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// errorFromChallenges returns the error from the given "WWW-Authenticate" header challenges.
|
||||
// It only looks at challenges with the "Bearer" scheme.
|
||||
func errorFromChallenges(cs []oauthex.Challenge) string {
|
||||
for _, c := range cs {
|
||||
if c.Scheme == "bearer" && c.Params["error"] != "" {
|
||||
return c.Params["error"]
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// getProtectedResourceMetadata returns the protected resource metadata.
|
||||
// If no metadata was found or the fetched metadata fails security checks,
|
||||
// it returns an error.
|
||||
func (h *AuthorizationCodeHandler) getProtectedResourceMetadata(ctx context.Context, wwwChallenges []oauthex.Challenge, mcpServerURL string) (*oauthex.ProtectedResourceMetadata, error) {
|
||||
var errs []error
|
||||
// Use MCP server URL as the resource URI per
|
||||
// https://modelcontextprotocol.io/specification/2025-11-25/basic/authorization#canonical-server-uri.
|
||||
for _, url := range protectedResourceMetadataURLs(resourceMetadataURLFromChallenges(wwwChallenges), mcpServerURL) {
|
||||
prm, err := oauthex.GetProtectedResourceMetadata(ctx, url.URL, url.Resource, h.config.Client)
|
||||
if err != nil {
|
||||
errs = append(errs, err)
|
||||
continue
|
||||
}
|
||||
if prm == nil {
|
||||
errs = append(errs, fmt.Errorf("protected resource metadata is nil"))
|
||||
continue
|
||||
}
|
||||
return prm, nil
|
||||
}
|
||||
return nil, fmt.Errorf("failed to get protected resource metadata: %v", errors.Join(errs...))
|
||||
}
|
||||
|
||||
type prmURL struct {
|
||||
// URL represents a URL where Protected Resource Metadata may be retrieved.
|
||||
URL string
|
||||
// Resource represents the corresponding resource URL for [URL].
|
||||
// It is required to perform validation described in RFC 9728, section 3.3.
|
||||
Resource string
|
||||
}
|
||||
|
||||
// protectedResourceMetadataURLs returns a list of URLs to try when looking for
|
||||
// protected resource metadata as mandated by the MCP specification:
|
||||
// https://modelcontextprotocol.io/specification/2025-11-25/basic/authorization#protected-resource-metadata-discovery-requirements
|
||||
func protectedResourceMetadataURLs(metadataURL, resourceURL string) []prmURL {
|
||||
var urls []prmURL
|
||||
if metadataURL != "" {
|
||||
urls = append(urls, prmURL{
|
||||
URL: metadataURL,
|
||||
Resource: resourceURL,
|
||||
})
|
||||
}
|
||||
ru, err := url.Parse(resourceURL)
|
||||
if err != nil {
|
||||
return urls
|
||||
}
|
||||
mu := *ru
|
||||
// "At the path of the server's MCP endpoint".
|
||||
mu.Path = "/.well-known/oauth-protected-resource/" + strings.TrimLeft(ru.Path, "/")
|
||||
urls = append(urls, prmURL{
|
||||
URL: mu.String(),
|
||||
Resource: resourceURL,
|
||||
})
|
||||
// "At the root".
|
||||
mu.Path = "/.well-known/oauth-protected-resource"
|
||||
ru.Path = ""
|
||||
urls = append(urls, prmURL{
|
||||
URL: mu.String(),
|
||||
Resource: ru.String(),
|
||||
})
|
||||
return urls
|
||||
}
|
||||
|
||||
// getAuthServerMetadata returns the authorization server metadata.
|
||||
// The provided Protected Resource Metadata must not be nil.
|
||||
// It returns an error if the metadata request fails with non-4xx HTTP status code
|
||||
// or the fetched metadata fails security checks.
|
||||
// If no metadata was found, it returns a minimal set of endpoints
|
||||
// as a fallback to 2025-03-26 spec.
|
||||
func (h *AuthorizationCodeHandler) getAuthServerMetadata(ctx context.Context, prm *oauthex.ProtectedResourceMetadata) (*oauthex.AuthServerMeta, error) {
|
||||
var authServerURL string
|
||||
if len(prm.AuthorizationServers) > 0 {
|
||||
// Use the first authorization server, similarly to other SDKs.
|
||||
authServerURL = prm.AuthorizationServers[0]
|
||||
} else {
|
||||
// Fallback to 2025-03-26 spec: MCP server base URL acts as Authorization Server.
|
||||
authURL, err := url.Parse(prm.Resource)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse resource URL: %v", err)
|
||||
}
|
||||
authURL.Path = ""
|
||||
authServerURL = authURL.String()
|
||||
}
|
||||
|
||||
for _, u := range authorizationServerMetadataURLs(authServerURL) {
|
||||
asm, err := oauthex.GetAuthServerMeta(ctx, u, authServerURL, h.config.Client)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get authorization server metadata: %w", err)
|
||||
}
|
||||
if asm != nil {
|
||||
return asm, nil
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback to 2025-03-26 spec: predefined endpoints.
|
||||
// https://modelcontextprotocol.io/specification/2025-03-26/basic/authorization#fallbacks-for-servers-without-metadata-discovery
|
||||
asm := &oauthex.AuthServerMeta{
|
||||
Issuer: authServerURL,
|
||||
AuthorizationEndpoint: authServerURL + "/authorize",
|
||||
TokenEndpoint: authServerURL + "/token",
|
||||
RegistrationEndpoint: authServerURL + "/register",
|
||||
}
|
||||
return asm, nil
|
||||
}
|
||||
|
||||
// authorizationServerMetadataURLs returns a list of URLs to try when looking for
|
||||
// authorization server metadata as mandated by the MCP specification:
|
||||
// https://modelcontextprotocol.io/specification/2025-11-25/basic/authorization#authorization-server-metadata-discovery.
|
||||
func authorizationServerMetadataURLs(issuerURL string) []string {
|
||||
var urls []string
|
||||
|
||||
baseURL, err := url.Parse(issuerURL)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if baseURL.Path == "" {
|
||||
// "OAuth 2.0 Authorization Server Metadata".
|
||||
baseURL.Path = "/.well-known/oauth-authorization-server"
|
||||
urls = append(urls, baseURL.String())
|
||||
// "OpenID Connect Discovery 1.0".
|
||||
baseURL.Path = "/.well-known/openid-configuration"
|
||||
urls = append(urls, baseURL.String())
|
||||
return urls
|
||||
}
|
||||
|
||||
originalPath := baseURL.Path
|
||||
// "OAuth 2.0 Authorization Server Metadata with path insertion".
|
||||
baseURL.Path = "/.well-known/oauth-authorization-server/" + strings.TrimLeft(originalPath, "/")
|
||||
urls = append(urls, baseURL.String())
|
||||
// "OpenID Connect Discovery 1.0 with path insertion".
|
||||
baseURL.Path = "/.well-known/openid-configuration/" + strings.TrimLeft(originalPath, "/")
|
||||
urls = append(urls, baseURL.String())
|
||||
// "OpenID Connect Discovery 1.0 with path appending".
|
||||
baseURL.Path = "/" + strings.Trim(originalPath, "/") + "/.well-known/openid-configuration"
|
||||
urls = append(urls, baseURL.String())
|
||||
return urls
|
||||
}
|
||||
|
||||
type registrationType int
|
||||
|
||||
const (
|
||||
registrationTypeClientIDMetadataDocument registrationType = iota
|
||||
registrationTypePreregistered
|
||||
registrationTypeDynamic
|
||||
)
|
||||
|
||||
type resolvedClientConfig struct {
|
||||
registrationType registrationType
|
||||
clientID string
|
||||
clientSecret string
|
||||
authStyle oauth2.AuthStyle
|
||||
}
|
||||
|
||||
func selectTokenAuthMethod(supported []string) oauth2.AuthStyle {
|
||||
prefOrder := []string{
|
||||
// Preferred in OAuth 2.1 draft: https://www.ietf.org/archive/id/draft-ietf-oauth-v2-1-14.html#name-client-secret.
|
||||
"client_secret_post",
|
||||
"client_secret_basic",
|
||||
}
|
||||
for _, method := range prefOrder {
|
||||
if slices.Contains(supported, method) {
|
||||
return authMethodToStyle(method)
|
||||
}
|
||||
}
|
||||
return oauth2.AuthStyleAutoDetect
|
||||
}
|
||||
|
||||
func authMethodToStyle(method string) oauth2.AuthStyle {
|
||||
switch method {
|
||||
case "client_secret_post":
|
||||
return oauth2.AuthStyleInParams
|
||||
case "client_secret_basic":
|
||||
return oauth2.AuthStyleInHeader
|
||||
case "none":
|
||||
// "none" is equivalent to "client_secret_post" but without sending client secret.
|
||||
return oauth2.AuthStyleInParams
|
||||
default:
|
||||
// "client_secret_basic" is the default per https://datatracker.ietf.org/doc/html/rfc7591#section-2.
|
||||
return oauth2.AuthStyleInHeader
|
||||
}
|
||||
}
|
||||
|
||||
// handleRegistration handles client registration.
|
||||
// The provided authorization server metadata must be non-nil.
|
||||
// Support for different registration methods is defined as follows:
|
||||
// - Client ID Metadata Document: metadata must have
|
||||
// `ClientIDMetadataDocumentSupported` set to true.
|
||||
// - Pre-registered client: assumed to be supported.
|
||||
// - Dynamic client registration: metadata must have
|
||||
// `RegistrationEndpoint` set to a non-empty value.
|
||||
func (h *AuthorizationCodeHandler) handleRegistration(ctx context.Context, asm *oauthex.AuthServerMeta) (*resolvedClientConfig, error) {
|
||||
// 1. Attempt to use Client ID Metadata Document (SEP-991).
|
||||
cimdCfg := h.config.ClientIDMetadataDocumentConfig
|
||||
if cimdCfg != nil && asm.ClientIDMetadataDocumentSupported {
|
||||
return &resolvedClientConfig{
|
||||
registrationType: registrationTypeClientIDMetadataDocument,
|
||||
clientID: cimdCfg.URL,
|
||||
}, nil
|
||||
}
|
||||
// 2. Attempt to use pre-registered client configuration.
|
||||
pCfg := h.config.PreregisteredClientConfig
|
||||
if pCfg != nil {
|
||||
authStyle := selectTokenAuthMethod(asm.TokenEndpointAuthMethodsSupported)
|
||||
return &resolvedClientConfig{
|
||||
registrationType: registrationTypePreregistered,
|
||||
clientID: pCfg.ClientSecretAuthConfig.ClientID,
|
||||
clientSecret: pCfg.ClientSecretAuthConfig.ClientSecret,
|
||||
authStyle: authStyle,
|
||||
}, nil
|
||||
}
|
||||
// 3. Attempt to use dynamic client registration.
|
||||
dcrCfg := h.config.DynamicClientRegistrationConfig
|
||||
if dcrCfg != nil && asm.RegistrationEndpoint != "" {
|
||||
regResp, err := oauthex.RegisterClient(ctx, asm.RegistrationEndpoint, dcrCfg.Metadata, h.config.Client)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to register client: %w", err)
|
||||
}
|
||||
cfg := &resolvedClientConfig{
|
||||
registrationType: registrationTypeDynamic,
|
||||
clientID: regResp.ClientID,
|
||||
clientSecret: regResp.ClientSecret,
|
||||
authStyle: authMethodToStyle(regResp.TokenEndpointAuthMethod),
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
return nil, fmt.Errorf("no configured client registration methods are supported by the authorization server")
|
||||
}
|
||||
|
||||
type authResult struct {
|
||||
*AuthorizationResult
|
||||
// usedCodeVerifier is the PKCE code verifier used to obtain the authorization code.
|
||||
// It is preserved for the token exchange step.
|
||||
usedCodeVerifier string
|
||||
}
|
||||
|
||||
// getAuthorizationCode uses the [AuthorizationCodeHandler.AuthorizationCodeFetcher]
|
||||
// to obtain an authorization code.
|
||||
func (h *AuthorizationCodeHandler) getAuthorizationCode(ctx context.Context, cfg *oauth2.Config, resourceURL string) (*authResult, error) {
|
||||
codeVerifier := oauth2.GenerateVerifier()
|
||||
state := rand.Text()
|
||||
|
||||
authURL := cfg.AuthCodeURL(state,
|
||||
oauth2.S256ChallengeOption(codeVerifier),
|
||||
oauth2.SetAuthURLParam("resource", resourceURL),
|
||||
)
|
||||
|
||||
authRes, err := h.config.AuthorizationCodeFetcher(ctx, &AuthorizationArgs{URL: authURL})
|
||||
if err != nil {
|
||||
// Purposefully leaving the error unwrappable so it can be handled by the caller.
|
||||
return nil, err
|
||||
}
|
||||
if authRes.State != state {
|
||||
return nil, fmt.Errorf("state mismatch")
|
||||
}
|
||||
return &authResult{
|
||||
AuthorizationResult: authRes,
|
||||
usedCodeVerifier: codeVerifier,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// exchangeAuthorizationCode exchanges the authorization code for a token
|
||||
// and stores it in a token source.
|
||||
func (h *AuthorizationCodeHandler) exchangeAuthorizationCode(ctx context.Context, cfg *oauth2.Config, authResult *authResult, resourceURL string) error {
|
||||
opts := []oauth2.AuthCodeOption{
|
||||
oauth2.VerifierOption(authResult.usedCodeVerifier),
|
||||
oauth2.SetAuthURLParam("resource", resourceURL),
|
||||
}
|
||||
clientCtx := context.WithValue(ctx, oauth2.HTTPClient, h.config.Client)
|
||||
token, err := cfg.Exchange(clientCtx, authResult.Code, opts...)
|
||||
if err != nil {
|
||||
return fmt.Errorf("token exchange failed: %w", err)
|
||||
}
|
||||
h.tokenSource = cfg.TokenSource(clientCtx, token)
|
||||
return nil
|
||||
}
|
||||
+42
@@ -0,0 +1,42 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
// OAuthHandler is an interface for handling OAuth flows.
|
||||
//
|
||||
// If a transport wishes to support OAuth 2 authorization, it should support
|
||||
// being configured with an OAuthHandler. It should call the handler's
|
||||
// TokenSource method whenever it sends an HTTP request to set the
|
||||
// Authorization header. If a request fails with a 401 or 403, it should call
|
||||
// Authorize, and if that returns nil, it should retry the request. It should
|
||||
// not call Authorize after the second failure. See
|
||||
// [github.com/modelcontextprotocol/go-sdk/mcp.StreamableClientTransport]
|
||||
// for an example.
|
||||
type OAuthHandler interface {
|
||||
isOAuthHandler()
|
||||
|
||||
// TokenSource returns a token source to be used for outgoing requests.
|
||||
// Returned token source might be nil. In that case, the transport will not
|
||||
// add any authorization headers to the request.
|
||||
TokenSource(context.Context) (oauth2.TokenSource, error)
|
||||
|
||||
// Authorize is called when an HTTP request results in an error that may
|
||||
// be addressed by the authorization flow (currently 401 Unauthorized and 403 Forbidden).
|
||||
// It is responsible for performing the OAuth flow to obtain an access token.
|
||||
// The arguments are the request that failed and the response that was received for it.
|
||||
// The headers of the request are available, but the body will have already been consumed
|
||||
// when Authorize is called.
|
||||
// If the returned error is nil, TokenSource is expected to return a non-nil token source.
|
||||
// After a successful call to Authorize, the HTTP request will be retried by the transport.
|
||||
// The function is responsible for closing the response body.
|
||||
Authorize(context.Context, *http.Request, *http.Response) error
|
||||
}
|
||||
+135
@@ -0,0 +1,135 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
//go:build mcp_go_client_oauth
|
||||
|
||||
package auth
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"sync"
|
||||
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
// An OAuthHandlerLegacy conducts an OAuth flow and returns a [oauth2.TokenSource] if the authorization
|
||||
// is approved, or an error if not.
|
||||
// The handler receives the HTTP request and response that triggered the authentication flow.
|
||||
// To obtain the protected resource metadata, call [oauthex.GetProtectedResourceMetadataFromHeader].
|
||||
//
|
||||
// Deprecated: Please use the new [OAuthHandler] abstraction that is built
|
||||
// into the streamable transport. This struct will be removed in v1.5.0.
|
||||
type OAuthHandlerLegacy func(req *http.Request, res *http.Response) (oauth2.TokenSource, error)
|
||||
|
||||
// HTTPTransport is an [http.RoundTripper] that follows the MCP
|
||||
// OAuth protocol when it encounters a 401 Unauthorized response.
|
||||
//
|
||||
// Deprecated: Please use the new [OAuthHandler] abstraction that is built
|
||||
// into the streamable transport. This struct will be removed in v1.5.0.
|
||||
type HTTPTransport struct {
|
||||
handler OAuthHandlerLegacy
|
||||
mu sync.Mutex // protects opts.Base
|
||||
opts HTTPTransportOptions
|
||||
}
|
||||
|
||||
// NewHTTPTransport returns a new [*HTTPTransport].
|
||||
// The handler is invoked when an HTTP request results in a 401 Unauthorized status.
|
||||
// It is called only once per transport. Once a TokenSource is obtained, it is used
|
||||
// for the lifetime of the transport; subsequent 401s are not processed.
|
||||
//
|
||||
// Deprecated: Please use the new [OAuthHandler] abstraction that is built
|
||||
// into the streamable transport. This struct will be removed in v1.5.0.
|
||||
func NewHTTPTransport(handler OAuthHandlerLegacy, opts *HTTPTransportOptions) (*HTTPTransport, error) {
|
||||
if handler == nil {
|
||||
return nil, errors.New("handler cannot be nil")
|
||||
}
|
||||
t := &HTTPTransport{
|
||||
handler: handler,
|
||||
}
|
||||
if opts != nil {
|
||||
t.opts = *opts
|
||||
}
|
||||
if t.opts.Base == nil {
|
||||
t.opts.Base = http.DefaultTransport
|
||||
}
|
||||
return t, nil
|
||||
}
|
||||
|
||||
// HTTPTransportOptions are options to [NewHTTPTransport].
|
||||
//
|
||||
// Deprecated: Please use the new [OAuthHandler] abstraction that is built
|
||||
// into the streamable transport. This struct will be removed in v1.5.0.
|
||||
type HTTPTransportOptions struct {
|
||||
// Base is the [http.RoundTripper] to use.
|
||||
// If nil, [http.DefaultTransport] is used.
|
||||
Base http.RoundTripper
|
||||
}
|
||||
|
||||
func (t *HTTPTransport) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
t.mu.Lock()
|
||||
base := t.opts.Base
|
||||
t.mu.Unlock()
|
||||
|
||||
var (
|
||||
// If haveBody is set, the request has a nontrivial body, and we need avoid
|
||||
// reading (or closing) it multiple times. In that case, bodyBytes is its
|
||||
// content.
|
||||
haveBody bool
|
||||
bodyBytes []byte
|
||||
)
|
||||
if req.Body != nil && req.Body != http.NoBody {
|
||||
// if we're setting Body, we must mutate first.
|
||||
req = req.Clone(req.Context())
|
||||
haveBody = true
|
||||
var err error
|
||||
bodyBytes, err = io.ReadAll(req.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Now that we've read the request body, http.RoundTripper requires that we
|
||||
// close it.
|
||||
req.Body.Close() // ignore error
|
||||
req.Body = io.NopCloser(bytes.NewReader(bodyBytes))
|
||||
}
|
||||
|
||||
resp, err := base.RoundTrip(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if resp.StatusCode != http.StatusUnauthorized {
|
||||
return resp, nil
|
||||
}
|
||||
if _, ok := base.(*oauth2.Transport); ok {
|
||||
// We failed to authorize even with a token source; give up.
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
resp.Body.Close()
|
||||
// Try to authorize.
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
// If we don't have a token source, get one by following the OAuth flow.
|
||||
// (We may have obtained one while t.mu was not held above.)
|
||||
// TODO: We hold the lock for the entire OAuth flow. This could be a long
|
||||
// time. Is there a better way?
|
||||
if _, ok := t.opts.Base.(*oauth2.Transport); !ok {
|
||||
ts, err := t.handler(req, resp)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
t.opts.Base = &oauth2.Transport{Base: t.opts.Base, Source: ts}
|
||||
}
|
||||
|
||||
// If we don't have a body, the request is reusable, though it will be cloned
|
||||
// by the base. However, if we've had to read the body, we must clone.
|
||||
if haveBody {
|
||||
req = req.Clone(req.Context())
|
||||
req.Body = io.NopCloser(bytes.NewReader(bodyBytes))
|
||||
}
|
||||
|
||||
return t.opts.Base.RoundTrip(req)
|
||||
}
|
||||
+19
@@ -0,0 +1,19 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by the license
|
||||
// that can be found in the LICENSE file.
|
||||
|
||||
// Package json provides internal JSON utilities.
|
||||
|
||||
package json
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
|
||||
"github.com/segmentio/encoding/json"
|
||||
)
|
||||
|
||||
func Unmarshal(data []byte, v any) error {
|
||||
dec := json.NewDecoder(bytes.NewReader(data))
|
||||
dec.DontMatchCaseInsensitiveStructFields()
|
||||
return dec.Decode(v)
|
||||
}
|
||||
+842
@@ -0,0 +1,842 @@
|
||||
// Copyright 2018 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package jsonrpc2
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/internal/json"
|
||||
)
|
||||
|
||||
// Binder builds a connection configuration.
|
||||
// This may be used in servers to generate a new configuration per connection.
|
||||
// ConnectionOptions itself implements Binder returning itself unmodified, to
|
||||
// allow for the simple cases where no per connection information is needed.
|
||||
type Binder interface {
|
||||
// Bind returns the ConnectionOptions to use when establishing the passed-in
|
||||
// Connection.
|
||||
//
|
||||
// The connection is not ready to use when Bind is called,
|
||||
// but Bind may close it without reading or writing to it.
|
||||
Bind(context.Context, *Connection) ConnectionOptions
|
||||
}
|
||||
|
||||
// A BinderFunc implements the Binder interface for a standalone Bind function.
|
||||
type BinderFunc func(context.Context, *Connection) ConnectionOptions
|
||||
|
||||
func (f BinderFunc) Bind(ctx context.Context, c *Connection) ConnectionOptions {
|
||||
return f(ctx, c)
|
||||
}
|
||||
|
||||
var _ Binder = BinderFunc(nil)
|
||||
|
||||
// ConnectionOptions holds the options for new connections.
|
||||
type ConnectionOptions struct {
|
||||
// Framer allows control over the message framing and encoding.
|
||||
// If nil, HeaderFramer will be used.
|
||||
Framer Framer
|
||||
// Preempter allows registration of a pre-queue message handler.
|
||||
// If nil, no messages will be preempted.
|
||||
Preempter Preempter
|
||||
// Handler is used as the queued message handler for inbound messages.
|
||||
// If nil, all responses will be ErrNotHandled.
|
||||
Handler Handler
|
||||
// OnInternalError, if non-nil, is called with any internal errors that occur
|
||||
// while serving the connection, such as protocol errors or invariant
|
||||
// violations. (If nil, internal errors result in panics.)
|
||||
OnInternalError func(error)
|
||||
}
|
||||
|
||||
// Connection manages the jsonrpc2 protocol, connecting responses back to their
|
||||
// calls. Connection is bidirectional; it does not have a designated server or
|
||||
// client end.
|
||||
//
|
||||
// Note that the word 'Connection' is overloaded: the mcp.Connection represents
|
||||
// the bidirectional stream of messages between client an server. The
|
||||
// jsonrpc2.Connection layers RPC logic on top of that stream, dispatching RPC
|
||||
// handlers, and correlating requests with responses from the peer.
|
||||
//
|
||||
// Some of the complexity of the Connection type is grown out of its usage in
|
||||
// gopls: it could probably be simplified based on our usage in MCP.
|
||||
type Connection struct {
|
||||
seq int64 // must only be accessed using atomic operations
|
||||
|
||||
stateMu sync.Mutex
|
||||
state inFlightState // accessed only in updateInFlight
|
||||
done chan struct{} // closed (under stateMu) when state.closed is true and all goroutines have completed
|
||||
|
||||
writer Writer
|
||||
handler Handler
|
||||
|
||||
onInternalError func(error)
|
||||
onDone func()
|
||||
}
|
||||
|
||||
// inFlightState records the state of the incoming and outgoing calls on a
|
||||
// Connection.
|
||||
type inFlightState struct {
|
||||
connClosing bool // true when the Connection's Close method has been called
|
||||
reading bool // true while the readIncoming goroutine is running
|
||||
readErr error // non-nil when the readIncoming goroutine exits (typically io.EOF)
|
||||
writeErr error // non-nil if a call to the Writer has failed with a non-canceled Context
|
||||
|
||||
// closer shuts down and cleans up the Reader and Writer state, ideally
|
||||
// interrupting any Read or Write call that is currently blocked. It is closed
|
||||
// when the state is idle and one of: connClosing is true, readErr is non-nil,
|
||||
// or writeErr is non-nil.
|
||||
//
|
||||
// After the closer has been invoked, the closer field is set to nil
|
||||
// and the closeErr field is simultaneously set to its result.
|
||||
closer io.Closer
|
||||
closeErr error // error returned from closer.Close
|
||||
|
||||
outgoingCalls map[ID]*AsyncCall // calls only
|
||||
outgoingNotifications int // # of notifications awaiting "write"
|
||||
|
||||
// incoming stores the total number of incoming calls and notifications
|
||||
// that have not yet written or processed a result.
|
||||
incoming int
|
||||
|
||||
incomingByID map[ID]*incomingRequest // calls only
|
||||
|
||||
// handlerQueue stores the backlog of calls and notifications that were not
|
||||
// already handled by a preempter.
|
||||
// The queue does not include the request currently being handled (if any).
|
||||
handlerQueue []*incomingRequest
|
||||
handlerRunning bool
|
||||
}
|
||||
|
||||
// updateInFlight locks the state of the connection's in-flight requests, allows
|
||||
// f to mutate that state, and closes the connection if it is idle and either
|
||||
// is closing or has a read or write error.
|
||||
func (c *Connection) updateInFlight(f func(*inFlightState)) {
|
||||
c.stateMu.Lock()
|
||||
defer c.stateMu.Unlock()
|
||||
|
||||
s := &c.state
|
||||
|
||||
f(s)
|
||||
|
||||
select {
|
||||
case <-c.done:
|
||||
// The connection was already completely done at the start of this call to
|
||||
// updateInFlight, so it must remain so. (The call to f should have noticed
|
||||
// that and avoided making any updates that would cause the state to be
|
||||
// non-idle.)
|
||||
if !s.idle() {
|
||||
panic("jsonrpc2: updateInFlight transitioned to non-idle when already done")
|
||||
}
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
if s.idle() && s.shuttingDown(ErrUnknown) != nil {
|
||||
if s.closer != nil {
|
||||
s.closeErr = s.closer.Close()
|
||||
s.closer = nil // prevent duplicate Close calls
|
||||
}
|
||||
if s.reading {
|
||||
// The readIncoming goroutine is still running. Our call to Close should
|
||||
// cause it to exit soon, at which point it will make another call to
|
||||
// updateInFlight, set s.reading to false, and mark the Connection done.
|
||||
} else {
|
||||
// The readIncoming goroutine has exited, or never started to begin with.
|
||||
// Since everything else is idle, we're completely done.
|
||||
if c.onDone != nil {
|
||||
c.onDone()
|
||||
}
|
||||
close(c.done)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// idle reports whether the connection is in a state with no pending calls or
|
||||
// notifications.
|
||||
//
|
||||
// If idle returns true, the readIncoming goroutine may still be running,
|
||||
// but no other goroutines are doing work on behalf of the connection.
|
||||
func (s *inFlightState) idle() bool {
|
||||
return len(s.outgoingCalls) == 0 && s.outgoingNotifications == 0 && s.incoming == 0 && !s.handlerRunning
|
||||
}
|
||||
|
||||
// shuttingDown reports whether the connection is in a state that should
|
||||
// disallow new (incoming and outgoing) calls. It returns either nil or
|
||||
// an error that is or wraps the provided errClosing.
|
||||
func (s *inFlightState) shuttingDown(errClosing error) error {
|
||||
if s.connClosing {
|
||||
// If Close has been called explicitly, it doesn't matter what state the
|
||||
// Reader and Writer are in: we shouldn't be starting new work because the
|
||||
// caller told us not to start new work.
|
||||
return errClosing
|
||||
}
|
||||
if s.readErr != nil {
|
||||
// If the read side of the connection is broken, we cannot read new call
|
||||
// requests, and cannot read responses to our outgoing calls.
|
||||
return fmt.Errorf("%w: %v", errClosing, s.readErr)
|
||||
}
|
||||
if s.writeErr != nil {
|
||||
// If the write side of the connection is broken, we cannot write responses
|
||||
// for incoming calls, and cannot write requests for outgoing calls.
|
||||
return fmt.Errorf("%w: %v", errClosing, s.writeErr)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// incomingRequest is used to track an incoming request as it is being handled
|
||||
type incomingRequest struct {
|
||||
*Request // the request being processed
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
// Bind returns the options unmodified.
|
||||
func (o ConnectionOptions) Bind(context.Context, *Connection) ConnectionOptions {
|
||||
return o
|
||||
}
|
||||
|
||||
// A ConnectionConfig configures a bidirectional jsonrpc2 connection.
|
||||
type ConnectionConfig struct {
|
||||
Reader Reader // required
|
||||
Writer Writer // required
|
||||
Closer io.Closer // required
|
||||
Preempter Preempter // optional
|
||||
Bind func(*Connection) Handler // required
|
||||
OnDone func() // optional
|
||||
OnInternalError func(error) // optional
|
||||
}
|
||||
|
||||
// NewConnection creates a new [Connection] object and starts processing
|
||||
// incoming messages.
|
||||
func NewConnection(ctx context.Context, cfg ConnectionConfig) *Connection {
|
||||
ctx = notDone{ctx}
|
||||
|
||||
c := &Connection{
|
||||
state: inFlightState{closer: cfg.Closer},
|
||||
done: make(chan struct{}),
|
||||
writer: cfg.Writer,
|
||||
onDone: cfg.OnDone,
|
||||
onInternalError: cfg.OnInternalError,
|
||||
}
|
||||
c.handler = cfg.Bind(c)
|
||||
c.start(ctx, cfg.Reader, cfg.Preempter)
|
||||
return c
|
||||
}
|
||||
|
||||
// bindConnection creates a new connection and runs it.
|
||||
//
|
||||
// This is used by the Dial and Serve functions to build the actual connection.
|
||||
//
|
||||
// The connection is closed automatically (and its resources cleaned up) when
|
||||
// the last request has completed after the underlying ReadWriteCloser breaks,
|
||||
// but it may be stopped earlier by calling Close (for a clean shutdown).
|
||||
func bindConnection(bindCtx context.Context, rwc io.ReadWriteCloser, binder Binder, onDone func()) *Connection {
|
||||
// TODO: Should we create a new event span here?
|
||||
// This will propagate cancellation from ctx; should it?
|
||||
ctx := notDone{bindCtx}
|
||||
|
||||
c := &Connection{
|
||||
state: inFlightState{closer: rwc},
|
||||
done: make(chan struct{}),
|
||||
onDone: onDone,
|
||||
}
|
||||
// It's tempting to set a finalizer on c to verify that the state has gone
|
||||
// idle when the connection becomes unreachable. Unfortunately, the Binder
|
||||
// interface makes that unsafe: it allows the Handler to close over the
|
||||
// Connection, which could create a reference cycle that would cause the
|
||||
// Connection to become uncollectable.
|
||||
|
||||
options := binder.Bind(bindCtx, c)
|
||||
framer := options.Framer
|
||||
if framer == nil {
|
||||
framer = HeaderFramer()
|
||||
}
|
||||
c.handler = options.Handler
|
||||
if c.handler == nil {
|
||||
c.handler = defaultHandler{}
|
||||
}
|
||||
c.onInternalError = options.OnInternalError
|
||||
|
||||
c.writer = framer.Writer(rwc)
|
||||
reader := framer.Reader(rwc)
|
||||
c.start(ctx, reader, options.Preempter)
|
||||
return c
|
||||
}
|
||||
|
||||
func (c *Connection) start(ctx context.Context, reader Reader, preempter Preempter) {
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
select {
|
||||
case <-c.done:
|
||||
// Bind already closed the connection; don't start a goroutine to read it.
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
// The goroutine started here will continue until the underlying stream is closed.
|
||||
//
|
||||
// (If the Binder closed the Connection already, this should error out and
|
||||
// return almost immediately.)
|
||||
s.reading = true
|
||||
go c.readIncoming(ctx, reader, preempter)
|
||||
})
|
||||
}
|
||||
|
||||
// Notify invokes the target method but does not wait for a response.
|
||||
// The params will be marshaled to JSON before sending over the wire, and will
|
||||
// be handed to the method invoked.
|
||||
func (c *Connection) Notify(ctx context.Context, method string, params any) (err error) {
|
||||
attempted := false
|
||||
|
||||
defer func() {
|
||||
if attempted {
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
s.outgoingNotifications--
|
||||
})
|
||||
}
|
||||
}()
|
||||
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
// If the connection is shutting down, allow outgoing notifications only if
|
||||
// there is at least one call still in flight. The number of calls in flight
|
||||
// cannot increase once shutdown begins, and allowing outgoing notifications
|
||||
// may permit notifications that will cancel in-flight calls.
|
||||
if len(s.outgoingCalls) == 0 && len(s.incomingByID) == 0 {
|
||||
err = s.shuttingDown(ErrClientClosing)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
s.outgoingNotifications++
|
||||
attempted = true
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
notify, err := NewNotification(method, params)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshaling notify parameters: %v", err)
|
||||
}
|
||||
|
||||
return c.write(ctx, notify)
|
||||
}
|
||||
|
||||
// Call invokes the target method and returns an object that can be used to await the response.
|
||||
// The params will be marshaled to JSON before sending over the wire, and will
|
||||
// be handed to the method invoked.
|
||||
// You do not have to wait for the response, it can just be ignored if not needed.
|
||||
// If sending the call failed, the response will be ready and have the error in it.
|
||||
func (c *Connection) Call(ctx context.Context, method string, params any) *AsyncCall {
|
||||
// Generate a new request identifier.
|
||||
id := Int64ID(atomic.AddInt64(&c.seq, 1))
|
||||
|
||||
ac := &AsyncCall{
|
||||
id: id,
|
||||
ready: make(chan struct{}),
|
||||
}
|
||||
// When this method returns, either ac is retired, or the request has been
|
||||
// written successfully and the call is awaiting a response (to be provided by
|
||||
// the readIncoming goroutine).
|
||||
|
||||
call, err := NewCall(ac.id, method, params)
|
||||
if err != nil {
|
||||
ac.retire(&Response{ID: id, Error: fmt.Errorf("marshaling call parameters: %w", err)})
|
||||
return ac
|
||||
}
|
||||
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
err = s.shuttingDown(ErrClientClosing)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if s.outgoingCalls == nil {
|
||||
s.outgoingCalls = make(map[ID]*AsyncCall)
|
||||
}
|
||||
s.outgoingCalls[ac.id] = ac
|
||||
})
|
||||
if err != nil {
|
||||
ac.retire(&Response{ID: id, Error: err})
|
||||
return ac
|
||||
}
|
||||
|
||||
if err := c.write(ctx, call); err != nil {
|
||||
// Sending failed. We will never get a response, so deliver a fake one if it
|
||||
// wasn't already retired by the connection breaking.
|
||||
c.Retire(ac, err)
|
||||
}
|
||||
return ac
|
||||
}
|
||||
|
||||
// Retire stops tracking the call, and reports err as its terminal error.
|
||||
//
|
||||
// Retire is safe to call multiple times: if the call is already no longer
|
||||
// tracked, Retire is a no op.
|
||||
func (c *Connection) Retire(ac *AsyncCall, err error) {
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
if s.outgoingCalls[ac.id] == ac {
|
||||
delete(s.outgoingCalls, ac.id)
|
||||
ac.retire(&Response{ID: ac.id, Error: err})
|
||||
} else {
|
||||
// ac was already retired elsewhere.
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// Async, signals that the current jsonrpc2 request may be handled
|
||||
// asynchronously to subsequent requests, when ctx is the request context.
|
||||
//
|
||||
// Async must be called at most once on each request's context (and its
|
||||
// descendants).
|
||||
func Async(ctx context.Context) {
|
||||
if r, ok := ctx.Value(asyncKey).(*releaser); ok {
|
||||
r.release(false)
|
||||
}
|
||||
}
|
||||
|
||||
type asyncKeyType struct{}
|
||||
|
||||
var asyncKey = asyncKeyType{}
|
||||
|
||||
// A releaser implements concurrency safe 'releasing' of async requests. (A
|
||||
// request is released when it is allowed to run concurrent with other
|
||||
// requests, via a call to [Async].)
|
||||
type releaser struct {
|
||||
mu sync.Mutex
|
||||
ch chan struct{}
|
||||
released bool
|
||||
}
|
||||
|
||||
// release closes the associated channel. If soft is set, multiple calls to
|
||||
// release are allowed.
|
||||
func (r *releaser) release(soft bool) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
if r.released {
|
||||
if !soft {
|
||||
panic("jsonrpc2.Async called multiple times")
|
||||
}
|
||||
} else {
|
||||
close(r.ch)
|
||||
r.released = true
|
||||
}
|
||||
}
|
||||
|
||||
type AsyncCall struct {
|
||||
id ID
|
||||
ready chan struct{} // closed after response has been set
|
||||
response *Response
|
||||
}
|
||||
|
||||
// ID used for this call.
|
||||
// This can be used to cancel the call if needed.
|
||||
func (ac *AsyncCall) ID() ID { return ac.id }
|
||||
|
||||
// IsReady can be used to check if the result is already prepared.
|
||||
// This is guaranteed to return true on a result for which Await has already
|
||||
// returned, or a call that failed to send in the first place.
|
||||
func (ac *AsyncCall) IsReady() bool {
|
||||
select {
|
||||
case <-ac.ready:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// retire processes the response to the call.
|
||||
//
|
||||
// It is an error to call retire more than once: retire is guarded by the
|
||||
// connection's outgoingCalls map.
|
||||
func (ac *AsyncCall) retire(response *Response) {
|
||||
select {
|
||||
case <-ac.ready:
|
||||
panic(fmt.Sprintf("jsonrpc2: retire called twice for ID %v", ac.id))
|
||||
default:
|
||||
}
|
||||
|
||||
ac.response = response
|
||||
close(ac.ready)
|
||||
}
|
||||
|
||||
// Await waits for (and decodes) the results of a Call.
|
||||
// The response will be unmarshaled from JSON into the result.
|
||||
//
|
||||
// If the call is cancelled due to context cancellation, the result is
|
||||
// ctx.Err().
|
||||
func (ac *AsyncCall) Await(ctx context.Context, result any) error {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-ac.ready:
|
||||
}
|
||||
if ac.response.Error != nil {
|
||||
return ac.response.Error
|
||||
}
|
||||
if result == nil {
|
||||
return nil
|
||||
}
|
||||
return json.Unmarshal(ac.response.Result, result)
|
||||
}
|
||||
|
||||
// Cancel cancels the Context passed to the Handle call for the inbound message
|
||||
// with the given ID.
|
||||
//
|
||||
// Cancel will not complain if the ID is not a currently active message, and it
|
||||
// will not cause any messages that have not arrived yet with that ID to be
|
||||
// cancelled.
|
||||
func (c *Connection) Cancel(id ID) {
|
||||
var req *incomingRequest
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
req = s.incomingByID[id]
|
||||
})
|
||||
if req != nil {
|
||||
req.cancel()
|
||||
}
|
||||
}
|
||||
|
||||
// Wait blocks until the connection is fully closed, but does not close it.
|
||||
func (c *Connection) Wait() error {
|
||||
return c.wait(true)
|
||||
}
|
||||
|
||||
// wait for the connection to close, and aggregates the most cause of its
|
||||
// termination, if abnormal.
|
||||
//
|
||||
// The fromWait argument allows this logic to be shared with Close, where we
|
||||
// only want to expose the closeErr.
|
||||
//
|
||||
// (Previously, Wait also only returned the closeErr, which was misleading if
|
||||
// the connection was broken for another reason).
|
||||
func (c *Connection) wait(fromWait bool) error {
|
||||
var err error
|
||||
<-c.done
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
if fromWait {
|
||||
if !errors.Is(s.readErr, io.EOF) {
|
||||
err = s.readErr
|
||||
}
|
||||
if err == nil && !errors.Is(s.writeErr, io.EOF) {
|
||||
err = s.writeErr
|
||||
}
|
||||
}
|
||||
if err == nil {
|
||||
err = s.closeErr
|
||||
}
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
// Close stops accepting new requests, waits for in-flight requests and enqueued
|
||||
// Handle calls to complete, and then closes the underlying stream.
|
||||
//
|
||||
// After the start of a Close, notification requests (that lack IDs and do not
|
||||
// receive responses) will continue to be passed to the Preempter, but calls
|
||||
// with IDs will receive immediate responses with ErrServerClosing, and no new
|
||||
// requests (not even notifications!) will be enqueued to the Handler.
|
||||
func (c *Connection) Close() error {
|
||||
// Stop handling new requests, and interrupt the reader (by closing the
|
||||
// connection) as soon as the active requests finish.
|
||||
c.updateInFlight(func(s *inFlightState) { s.connClosing = true })
|
||||
return c.wait(false)
|
||||
}
|
||||
|
||||
// readIncoming collects inbound messages from the reader and delivers them, either responding
|
||||
// to outgoing calls or feeding requests to the queue.
|
||||
func (c *Connection) readIncoming(ctx context.Context, reader Reader, preempter Preempter) {
|
||||
var err error
|
||||
for {
|
||||
var msg Message
|
||||
msg, err = reader.Read(ctx)
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
|
||||
switch msg := msg.(type) {
|
||||
case *Request:
|
||||
c.acceptRequest(ctx, msg, preempter)
|
||||
|
||||
case *Response:
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
if ac, ok := s.outgoingCalls[msg.ID]; ok {
|
||||
delete(s.outgoingCalls, msg.ID)
|
||||
ac.retire(msg)
|
||||
} else {
|
||||
// TODO: How should we report unexpected responses?
|
||||
}
|
||||
})
|
||||
|
||||
default:
|
||||
c.internalErrorf("Read returned an unexpected message of type %T", msg)
|
||||
}
|
||||
}
|
||||
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
s.reading = false
|
||||
s.readErr = err
|
||||
|
||||
// Retire any outgoing requests that were still in flight: with the Reader no
|
||||
// longer being processed, they necessarily cannot receive a response.
|
||||
for id, ac := range s.outgoingCalls {
|
||||
ac.retire(&Response{ID: id, Error: err})
|
||||
}
|
||||
s.outgoingCalls = nil
|
||||
})
|
||||
}
|
||||
|
||||
// acceptRequest either handles msg synchronously or enqueues it to be handled
|
||||
// asynchronously.
|
||||
func (c *Connection) acceptRequest(ctx context.Context, msg *Request, preempter Preempter) {
|
||||
// In theory notifications cannot be cancelled, but we build them a cancel
|
||||
// context anyway.
|
||||
reqCtx, cancel := context.WithCancel(ctx)
|
||||
req := &incomingRequest{
|
||||
Request: msg,
|
||||
ctx: reqCtx,
|
||||
cancel: cancel,
|
||||
}
|
||||
|
||||
// If the request is a call, add it to the incoming map so it can be
|
||||
// cancelled (or responded) by ID.
|
||||
var err error
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
s.incoming++
|
||||
|
||||
if req.IsCall() {
|
||||
if s.incomingByID[req.ID] != nil {
|
||||
err = fmt.Errorf("%w: request ID %v already in use", ErrInvalidRequest, req.ID)
|
||||
req.ID = ID{} // Don't misattribute this error to the existing request.
|
||||
return
|
||||
}
|
||||
|
||||
if s.incomingByID == nil {
|
||||
s.incomingByID = make(map[ID]*incomingRequest)
|
||||
}
|
||||
s.incomingByID[req.ID] = req
|
||||
|
||||
// When shutting down, reject all new Call requests, even if they could
|
||||
// theoretically be handled by the preempter. The preempter could return
|
||||
// ErrAsyncResponse, which would increase the amount of work in flight
|
||||
// when we're trying to ensure that it strictly decreases.
|
||||
err = s.shuttingDown(ErrServerClosing)
|
||||
}
|
||||
})
|
||||
if err != nil {
|
||||
c.processResult("acceptRequest", req, nil, err)
|
||||
return
|
||||
}
|
||||
|
||||
if preempter != nil {
|
||||
result, err := preempter.Preempt(req.ctx, req.Request)
|
||||
|
||||
if !errors.Is(err, ErrNotHandled) {
|
||||
c.processResult("Preempt", req, result, err)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
// If the connection is shutting down, don't enqueue anything to the
|
||||
// handler — not even notifications. That ensures that if the handler
|
||||
// continues to make progress, it will eventually become idle and
|
||||
// close the connection.
|
||||
err = s.shuttingDown(ErrServerClosing)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
// We enqueue requests that have not been preempted to an unbounded slice.
|
||||
// Unfortunately, we cannot in general limit the size of the handler
|
||||
// queue: we have to read every response that comes in on the wire
|
||||
// (because it may be responding to a request issued by, say, an
|
||||
// asynchronous handler), and in order to get to that response we have
|
||||
// to read all of the requests that came in ahead of it.
|
||||
s.handlerQueue = append(s.handlerQueue, req)
|
||||
if !s.handlerRunning {
|
||||
// We start the handleAsync goroutine when it has work to do, and let it
|
||||
// exit when the queue empties.
|
||||
//
|
||||
// Otherwise, in order to synchronize the handler we would need some other
|
||||
// goroutine (probably readIncoming?) to explicitly wait for handleAsync
|
||||
// to finish, and that would complicate error reporting: either the error
|
||||
// report from the goroutine would be blocked on the handler emptying its
|
||||
// queue (which was tried, and introduced a deadlock detected by
|
||||
// TestCloseCallRace), or the error would need to be reported separately
|
||||
// from synchronizing completion. Allowing the handler goroutine to exit
|
||||
// when idle seems simpler than trying to implement either of those
|
||||
// alternatives correctly.
|
||||
s.handlerRunning = true
|
||||
go c.handleAsync()
|
||||
}
|
||||
})
|
||||
if err != nil {
|
||||
c.processResult("acceptRequest", req, nil, err)
|
||||
}
|
||||
}
|
||||
|
||||
// handleAsync invokes the handler on the requests in the handler queue
|
||||
// sequentially until the queue is empty.
|
||||
func (c *Connection) handleAsync() {
|
||||
for {
|
||||
var req *incomingRequest
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
if len(s.handlerQueue) > 0 {
|
||||
req, s.handlerQueue = s.handlerQueue[0], s.handlerQueue[1:]
|
||||
} else {
|
||||
s.handlerRunning = false
|
||||
}
|
||||
})
|
||||
if req == nil {
|
||||
return
|
||||
}
|
||||
|
||||
// Only deliver to the Handler if not already canceled.
|
||||
if err := req.ctx.Err(); err != nil {
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
if s.writeErr != nil {
|
||||
// Assume that req.ctx was canceled due to s.writeErr.
|
||||
// TODO(#51365): use a Context API to plumb this through req.ctx.
|
||||
err = fmt.Errorf("%w: %v", ErrServerClosing, s.writeErr)
|
||||
}
|
||||
})
|
||||
c.processResult("handleAsync", req, nil, err)
|
||||
continue
|
||||
}
|
||||
|
||||
releaser := &releaser{ch: make(chan struct{})}
|
||||
ctx := context.WithValue(req.ctx, asyncKey, releaser)
|
||||
go func() {
|
||||
defer releaser.release(true)
|
||||
result, err := c.handler.Handle(ctx, req.Request)
|
||||
c.processResult(c.handler, req, result, err)
|
||||
}()
|
||||
<-releaser.ch
|
||||
}
|
||||
}
|
||||
|
||||
// processResult processes the result of a request and, if appropriate, sends a response.
|
||||
func (c *Connection) processResult(from any, req *incomingRequest, result any, err error) error {
|
||||
switch err {
|
||||
case ErrNotHandled, ErrMethodNotFound:
|
||||
// Add detail describing the unhandled method.
|
||||
err = fmt.Errorf("%w: %q", ErrMethodNotFound, req.Method)
|
||||
}
|
||||
|
||||
if result != nil && err != nil {
|
||||
c.internalErrorf("%#v returned a non-nil result with a non-nil error for %s:\n%v\n%#v", from, req.Method, err, result)
|
||||
result = nil // Discard the spurious result and respond with err.
|
||||
}
|
||||
|
||||
if req.IsCall() {
|
||||
if result == nil && err == nil {
|
||||
err = c.internalErrorf("%#v returned a nil result and nil error for a %q Request that requires a Response", from, req.Method)
|
||||
}
|
||||
|
||||
response, respErr := NewResponse(req.ID, result, err)
|
||||
|
||||
// The caller could theoretically reuse the request's ID as soon as we've
|
||||
// sent the response, so ensure that it is removed from the incoming map
|
||||
// before sending.
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
delete(s.incomingByID, req.ID)
|
||||
})
|
||||
if respErr == nil {
|
||||
writeErr := c.write(notDone{req.ctx}, response)
|
||||
if err == nil {
|
||||
err = writeErr
|
||||
}
|
||||
} else {
|
||||
err = c.internalErrorf("%#v returned a malformed result for %q: %w", from, req.Method, respErr)
|
||||
}
|
||||
} else { // req is a notification
|
||||
if result != nil {
|
||||
err = c.internalErrorf("%#v returned a non-nil result for a %q Request without an ID", from, req.Method)
|
||||
} else if err != nil {
|
||||
err = fmt.Errorf("%w: %q notification failed: %v", ErrInternal, req.Method, err)
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
// TODO: can/should we do anything with this error beyond writing it to the event log?
|
||||
// (Is this the right label to attach to the log?)
|
||||
}
|
||||
|
||||
// Cancel the request to free any associated resources.
|
||||
req.cancel()
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
if s.incoming == 0 {
|
||||
panic("jsonrpc2: processResult called when incoming count is already zero")
|
||||
}
|
||||
s.incoming--
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
// write is used by all things that write outgoing messages, including replies.
|
||||
// it makes sure that writes are atomic
|
||||
func (c *Connection) write(ctx context.Context, msg Message) error {
|
||||
var err error
|
||||
// Fail writes immediately if the connection is shutting down.
|
||||
//
|
||||
// TODO(rfindley): should we allow cancellation notifications through? It
|
||||
// could be the case that writes can still succeed.
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
err = s.shuttingDown(ErrServerClosing)
|
||||
})
|
||||
if err == nil {
|
||||
err = c.writer.Write(ctx, msg)
|
||||
}
|
||||
|
||||
// For cancelled or rejected requests, we don't set the writeErr (which would
|
||||
// break the connection). They can just be returned to the caller.
|
||||
if err != nil && ctx.Err() == nil && !errors.Is(err, ErrRejected) {
|
||||
// The call to Write failed, and since ctx.Err() is nil we can't attribute
|
||||
// the failure (even indirectly) to Context cancellation. The writer appears
|
||||
// to be broken, and future writes are likely to also fail.
|
||||
//
|
||||
// If the read side of the connection is also broken, we might not even be
|
||||
// able to receive cancellation notifications. Since we can't reliably write
|
||||
// the results of incoming calls and can't receive explicit cancellations,
|
||||
// cancel the calls now.
|
||||
c.updateInFlight(func(s *inFlightState) {
|
||||
if s.writeErr == nil {
|
||||
s.writeErr = err
|
||||
for _, r := range s.incomingByID {
|
||||
r.cancel()
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
// internalErrorf reports an internal error. By default it panics, but if
|
||||
// c.onInternalError is non-nil it instead calls that and returns an error
|
||||
// wrapping ErrInternal.
|
||||
func (c *Connection) internalErrorf(format string, args ...any) error {
|
||||
err := fmt.Errorf(format, args...)
|
||||
if c.onInternalError == nil {
|
||||
panic("jsonrpc2: " + err.Error())
|
||||
}
|
||||
c.onInternalError(err)
|
||||
|
||||
return fmt.Errorf("%w: %v", ErrInternal, err)
|
||||
}
|
||||
|
||||
// notDone is a context.Context wrapper that returns a nil Done channel.
|
||||
type notDone struct{ ctx context.Context }
|
||||
|
||||
func (ic notDone) Value(key any) any {
|
||||
return ic.ctx.Value(key)
|
||||
}
|
||||
|
||||
func (notDone) Done() <-chan struct{} { return nil }
|
||||
func (notDone) Err() error { return nil }
|
||||
func (notDone) Deadline() (time.Time, bool) { return time.Time{}, false }
|
||||
+208
@@ -0,0 +1,208 @@
|
||||
// Copyright 2018 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package jsonrpc2
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// Reader abstracts the transport mechanics from the JSON RPC protocol.
|
||||
// A Conn reads messages from the reader it was provided on construction,
|
||||
// and assumes that each call to Read fully transfers a single message,
|
||||
// or returns an error.
|
||||
//
|
||||
// A reader is not safe for concurrent use, it is expected it will be used by
|
||||
// a single Conn in a safe manner.
|
||||
type Reader interface {
|
||||
// Read gets the next message from the stream.
|
||||
Read(context.Context) (Message, error)
|
||||
}
|
||||
|
||||
// Writer abstracts the transport mechanics from the JSON RPC protocol.
|
||||
// A Conn writes messages using the writer it was provided on construction,
|
||||
// and assumes that each call to Write fully transfers a single message,
|
||||
// or returns an error.
|
||||
//
|
||||
// A writer must be safe for concurrent use, as writes may occur concurrently
|
||||
// in practice: libraries may make calls or respond to requests asynchronously.
|
||||
type Writer interface {
|
||||
// Write sends a message to the stream.
|
||||
Write(context.Context, Message) error
|
||||
}
|
||||
|
||||
// Framer wraps low level byte readers and writers into jsonrpc2 message
|
||||
// readers and writers.
|
||||
// It is responsible for the framing and encoding of messages into wire form.
|
||||
//
|
||||
// TODO(rfindley): rethink the framer interface, as with JSONRPC2 batching
|
||||
// there is a need for Reader and Writer to be correlated, and while the
|
||||
// implementation of framing here allows that, it is not made explicit by the
|
||||
// interface.
|
||||
//
|
||||
// Perhaps a better interface would be
|
||||
//
|
||||
// Frame(io.ReadWriteCloser) (Reader, Writer).
|
||||
type Framer interface {
|
||||
// Reader wraps a byte reader into a message reader.
|
||||
Reader(io.Reader) Reader
|
||||
// Writer wraps a byte writer into a message writer.
|
||||
Writer(io.Writer) Writer
|
||||
}
|
||||
|
||||
// RawFramer returns a new Framer.
|
||||
// The messages are sent with no wrapping, and rely on json decode consistency
|
||||
// to determine message boundaries.
|
||||
func RawFramer() Framer { return rawFramer{} }
|
||||
|
||||
type rawFramer struct{}
|
||||
type rawReader struct{ in *json.Decoder }
|
||||
type rawWriter struct {
|
||||
mu sync.Mutex
|
||||
out io.Writer
|
||||
}
|
||||
|
||||
func (rawFramer) Reader(rw io.Reader) Reader {
|
||||
return &rawReader{in: json.NewDecoder(rw)}
|
||||
}
|
||||
|
||||
func (rawFramer) Writer(rw io.Writer) Writer {
|
||||
return &rawWriter{out: rw}
|
||||
}
|
||||
|
||||
func (r *rawReader) Read(ctx context.Context) (Message, error) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
default:
|
||||
}
|
||||
var raw json.RawMessage
|
||||
if err := r.in.Decode(&raw); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
msg, err := DecodeMessage(raw)
|
||||
return msg, err
|
||||
}
|
||||
|
||||
func (w *rawWriter) Write(ctx context.Context, msg Message) error {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
default:
|
||||
}
|
||||
|
||||
data, err := EncodeMessage(msg)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshaling message: %v", err)
|
||||
}
|
||||
|
||||
w.mu.Lock()
|
||||
defer w.mu.Unlock()
|
||||
_, err = w.out.Write(data)
|
||||
return err
|
||||
}
|
||||
|
||||
// HeaderFramer returns a new Framer.
|
||||
// The messages are sent with HTTP content length and MIME type headers.
|
||||
// This is the format used by LSP and others.
|
||||
func HeaderFramer() Framer { return headerFramer{} }
|
||||
|
||||
type headerFramer struct{}
|
||||
type headerReader struct{ in *bufio.Reader }
|
||||
type headerWriter struct {
|
||||
mu sync.Mutex
|
||||
out io.Writer
|
||||
}
|
||||
|
||||
func (headerFramer) Reader(rw io.Reader) Reader {
|
||||
return &headerReader{in: bufio.NewReader(rw)}
|
||||
}
|
||||
|
||||
func (headerFramer) Writer(rw io.Writer) Writer {
|
||||
return &headerWriter{out: rw}
|
||||
}
|
||||
|
||||
func (r *headerReader) Read(ctx context.Context) (Message, error) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
default:
|
||||
}
|
||||
|
||||
firstRead := true // to detect a clean EOF below
|
||||
var contentLength int64
|
||||
// read the header, stop on the first empty line
|
||||
for {
|
||||
line, err := r.in.ReadString('\n')
|
||||
if err != nil {
|
||||
if err == io.EOF {
|
||||
if firstRead && line == "" {
|
||||
return nil, io.EOF // clean EOF
|
||||
}
|
||||
err = io.ErrUnexpectedEOF
|
||||
}
|
||||
return nil, fmt.Errorf("failed reading header line: %w", err)
|
||||
}
|
||||
firstRead = false
|
||||
|
||||
line = strings.TrimSpace(line)
|
||||
// check we have a header line
|
||||
if line == "" {
|
||||
break
|
||||
}
|
||||
colon := strings.IndexRune(line, ':')
|
||||
if colon < 0 {
|
||||
return nil, fmt.Errorf("invalid header line %q", line)
|
||||
}
|
||||
name, value := line[:colon], strings.TrimSpace(line[colon+1:])
|
||||
switch {
|
||||
case strings.EqualFold(name, "Content-Length"):
|
||||
if contentLength, err = strconv.ParseInt(value, 10, 32); err != nil {
|
||||
return nil, fmt.Errorf("failed parsing Content-Length: %v", value)
|
||||
}
|
||||
if contentLength <= 0 {
|
||||
return nil, fmt.Errorf("invalid Content-Length: %v", contentLength)
|
||||
}
|
||||
default:
|
||||
// ignoring unknown headers
|
||||
}
|
||||
}
|
||||
if contentLength == 0 {
|
||||
return nil, fmt.Errorf("missing Content-Length header")
|
||||
}
|
||||
data := make([]byte, contentLength)
|
||||
_, err := io.ReadFull(r.in, data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
msg, err := DecodeMessage(data)
|
||||
return msg, err
|
||||
}
|
||||
|
||||
func (w *headerWriter) Write(ctx context.Context, msg Message) error {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
default:
|
||||
}
|
||||
w.mu.Lock()
|
||||
defer w.mu.Unlock()
|
||||
|
||||
data, err := EncodeMessage(msg)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshaling message: %v", err)
|
||||
}
|
||||
_, err = fmt.Fprintf(w.out, "Content-Length: %v\r\n\r\n", len(data))
|
||||
if err == nil {
|
||||
_, err = w.out.Write(data)
|
||||
}
|
||||
return err
|
||||
}
|
||||
+121
@@ -0,0 +1,121 @@
|
||||
// Copyright 2018 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// Package jsonrpc2 is a minimal implementation of the JSON RPC 2 spec.
|
||||
// https://www.jsonrpc.org/specification
|
||||
// It is intended to be compatible with other implementations at the wire level.
|
||||
package jsonrpc2
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
)
|
||||
|
||||
var (
|
||||
// ErrIdleTimeout is returned when serving timed out waiting for new connections.
|
||||
ErrIdleTimeout = errors.New("timed out waiting for new connections")
|
||||
|
||||
// ErrNotHandled is returned from a Handler or Preempter to indicate it did
|
||||
// not handle the request.
|
||||
//
|
||||
// If a Handler returns ErrNotHandled, the server replies with
|
||||
// ErrMethodNotFound.
|
||||
ErrNotHandled = errors.New("JSON RPC not handled")
|
||||
)
|
||||
|
||||
// Preempter handles messages on a connection before they are queued to the main
|
||||
// handler.
|
||||
// Primarily this is used for cancel handlers or notifications for which out of
|
||||
// order processing is not an issue.
|
||||
type Preempter interface {
|
||||
// Preempt is invoked for each incoming request before it is queued for handling.
|
||||
//
|
||||
// If Preempt returns ErrNotHandled, the request will be queued,
|
||||
// and eventually passed to a Handle call.
|
||||
//
|
||||
// Otherwise, the result and error are processed as if returned by Handle.
|
||||
//
|
||||
// Preempt must not block. (The Context passed to it is for Values only.)
|
||||
Preempt(ctx context.Context, req *Request) (result any, err error)
|
||||
}
|
||||
|
||||
// A PreempterFunc implements the Preempter interface for a standalone Preempt function.
|
||||
type PreempterFunc func(ctx context.Context, req *Request) (any, error)
|
||||
|
||||
func (f PreempterFunc) Preempt(ctx context.Context, req *Request) (any, error) {
|
||||
return f(ctx, req)
|
||||
}
|
||||
|
||||
var _ Preempter = PreempterFunc(nil)
|
||||
|
||||
// Handler handles messages on a connection.
|
||||
type Handler interface {
|
||||
// Handle is invoked sequentially for each incoming request that has not
|
||||
// already been handled by a Preempter.
|
||||
//
|
||||
// If the Request has a nil ID, Handle must return a nil result,
|
||||
// and any error may be logged but will not be reported to the caller.
|
||||
//
|
||||
// If the Request has a non-nil ID, Handle must return either a
|
||||
// non-nil, JSON-marshalable result, or a non-nil error.
|
||||
//
|
||||
// The Context passed to Handle will be canceled if the
|
||||
// connection is broken or the request is canceled or completed.
|
||||
// (If Handle returns ErrAsyncResponse, ctx will remain uncanceled
|
||||
// until either Cancel or Respond is called for the request's ID.)
|
||||
Handle(ctx context.Context, req *Request) (result any, err error)
|
||||
}
|
||||
|
||||
type defaultHandler struct{}
|
||||
|
||||
func (defaultHandler) Preempt(context.Context, *Request) (any, error) {
|
||||
return nil, ErrNotHandled
|
||||
}
|
||||
|
||||
func (defaultHandler) Handle(context.Context, *Request) (any, error) {
|
||||
return nil, ErrNotHandled
|
||||
}
|
||||
|
||||
// A HandlerFunc implements the Handler interface for a standalone Handle function.
|
||||
type HandlerFunc func(ctx context.Context, req *Request) (any, error)
|
||||
|
||||
func (f HandlerFunc) Handle(ctx context.Context, req *Request) (any, error) {
|
||||
return f(ctx, req)
|
||||
}
|
||||
|
||||
var _ Handler = HandlerFunc(nil)
|
||||
|
||||
// async is a small helper for operations with an asynchronous result that you
|
||||
// can wait for.
|
||||
type async struct {
|
||||
ready chan struct{} // closed when done
|
||||
firstErr chan error // 1-buffered; contains either nil or the first non-nil error
|
||||
}
|
||||
|
||||
func newAsync() *async {
|
||||
var a async
|
||||
a.ready = make(chan struct{})
|
||||
a.firstErr = make(chan error, 1)
|
||||
a.firstErr <- nil
|
||||
return &a
|
||||
}
|
||||
|
||||
func (a *async) done() {
|
||||
close(a.ready)
|
||||
}
|
||||
|
||||
func (a *async) wait() error {
|
||||
<-a.ready
|
||||
err := <-a.firstErr
|
||||
a.firstErr <- err
|
||||
return err
|
||||
}
|
||||
|
||||
func (a *async) setError(err error) {
|
||||
storedErr := <-a.firstErr
|
||||
if storedErr == nil {
|
||||
storedErr = err
|
||||
}
|
||||
a.firstErr <- storedErr
|
||||
}
|
||||
+242
@@ -0,0 +1,242 @@
|
||||
// Copyright 2018 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package jsonrpc2
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
internaljson "github.com/modelcontextprotocol/go-sdk/internal/json"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/internal/mcpgodebug"
|
||||
)
|
||||
|
||||
// ID is a Request identifier, which is defined by the spec to be a string, integer, or null.
|
||||
// https://www.jsonrpc.org/specification#request_object
|
||||
type ID struct {
|
||||
value any
|
||||
}
|
||||
|
||||
// MakeID coerces the given Go value to an ID. The value should be the
|
||||
// default JSON marshaling of a Request identifier: nil, float64, or string.
|
||||
//
|
||||
// Returns an error if the value type was not a valid Request ID type.
|
||||
//
|
||||
// TODO: ID can't be a json.Marshaler/Unmarshaler, because we want to omitzero.
|
||||
// Simplify this package by making ID json serializable once we can rely on
|
||||
// omitzero.
|
||||
func MakeID(v any) (ID, error) {
|
||||
switch v := v.(type) {
|
||||
case nil:
|
||||
return ID{}, nil
|
||||
case float64:
|
||||
return Int64ID(int64(v)), nil
|
||||
case string:
|
||||
return StringID(v), nil
|
||||
}
|
||||
return ID{}, fmt.Errorf("%w: invalid ID type %T", ErrParse, v)
|
||||
}
|
||||
|
||||
// Message is the interface to all jsonrpc2 message types.
|
||||
// They share no common functionality, but are a closed set of concrete types
|
||||
// that are allowed to implement this interface. The message types are *Request
|
||||
// and *Response.
|
||||
type Message interface {
|
||||
// marshal builds the wire form from the API form.
|
||||
// It is private, which makes the set of Message implementations closed.
|
||||
marshal(to *wireCombined)
|
||||
}
|
||||
|
||||
// Request is a Message sent to a peer to request behavior.
|
||||
// If it has an ID it is a call, otherwise it is a notification.
|
||||
type Request struct {
|
||||
// ID of this request, used to tie the Response back to the request.
|
||||
// This will be nil for notifications.
|
||||
ID ID
|
||||
// Method is a string containing the method name to invoke.
|
||||
Method string
|
||||
// Params is either a struct or an array with the parameters of the method.
|
||||
Params json.RawMessage
|
||||
// Extra is additional information that does not appear on the wire. It can be
|
||||
// used to pass information from the application to the underlying transport.
|
||||
Extra any
|
||||
}
|
||||
|
||||
// Response is a Message used as a reply to a call Request.
|
||||
// It will have the same ID as the call it is a response to.
|
||||
type Response struct {
|
||||
// result is the content of the response.
|
||||
Result json.RawMessage
|
||||
// err is set only if the call failed.
|
||||
Error error
|
||||
// id of the request this is a response to.
|
||||
ID ID
|
||||
// Extra is additional information that does not appear on the wire. It can be
|
||||
// used to pass information from the underlying transport to the application.
|
||||
Extra any
|
||||
}
|
||||
|
||||
// StringID creates a new string request identifier.
|
||||
func StringID(s string) ID { return ID{value: s} }
|
||||
|
||||
// Int64ID creates a new integer request identifier.
|
||||
func Int64ID(i int64) ID { return ID{value: i} }
|
||||
|
||||
// IsValid returns true if the ID is a valid identifier.
|
||||
// The default value for ID will return false.
|
||||
func (id ID) IsValid() bool { return id.value != nil }
|
||||
|
||||
// Raw returns the underlying value of the ID.
|
||||
func (id ID) Raw() any { return id.value }
|
||||
|
||||
// NewNotification constructs a new Notification message for the supplied
|
||||
// method and parameters.
|
||||
func NewNotification(method string, params any) (*Request, error) {
|
||||
p, merr := marshalToRaw(params)
|
||||
return &Request{Method: method, Params: p}, merr
|
||||
}
|
||||
|
||||
// NewCall constructs a new Call message for the supplied ID, method and
|
||||
// parameters.
|
||||
func NewCall(id ID, method string, params any) (*Request, error) {
|
||||
p, merr := marshalToRaw(params)
|
||||
return &Request{ID: id, Method: method, Params: p}, merr
|
||||
}
|
||||
|
||||
func (msg *Request) IsCall() bool { return msg.ID.IsValid() }
|
||||
|
||||
func (msg *Request) marshal(to *wireCombined) {
|
||||
to.ID = msg.ID.value
|
||||
to.Method = msg.Method
|
||||
to.Params = msg.Params
|
||||
}
|
||||
|
||||
// NewResponse constructs a new Response message that is a reply to the
|
||||
// supplied. If err is set result may be ignored.
|
||||
func NewResponse(id ID, result any, rerr error) (*Response, error) {
|
||||
r, merr := marshalToRaw(result)
|
||||
return &Response{ID: id, Result: r, Error: rerr}, merr
|
||||
}
|
||||
|
||||
func (msg *Response) marshal(to *wireCombined) {
|
||||
to.ID = msg.ID.value
|
||||
to.Error = toWireError(msg.Error)
|
||||
to.Result = msg.Result
|
||||
}
|
||||
|
||||
func toWireError(err error) *WireError {
|
||||
if err == nil {
|
||||
// no error, the response is complete
|
||||
return nil
|
||||
}
|
||||
if err, ok := err.(*WireError); ok {
|
||||
// already a wire error, just use it
|
||||
return err
|
||||
}
|
||||
result := &WireError{Message: err.Error()}
|
||||
var wrapped *WireError
|
||||
if errors.As(err, &wrapped) {
|
||||
// if we wrapped a wire error, keep the code from the wrapped error
|
||||
// but the message from the outer error
|
||||
result.Code = wrapped.Code
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func EncodeMessage(msg Message) ([]byte, error) {
|
||||
wire := wireCombined{VersionTag: wireVersion}
|
||||
msg.marshal(&wire)
|
||||
data, err := jsonMarshal(&wire)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshaling jsonrpc message: %w", err)
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// EncodeIndent is like EncodeMessage, but honors indents.
|
||||
// TODO(rfindley): refactor so that this concern is handled independently.
|
||||
// Perhaps we should pass in a json.Encoder?
|
||||
func EncodeIndent(msg Message, prefix, indent string) ([]byte, error) {
|
||||
wire := wireCombined{VersionTag: wireVersion}
|
||||
msg.marshal(&wire)
|
||||
var buf bytes.Buffer
|
||||
enc := json.NewEncoder(&buf)
|
||||
enc.SetEscapeHTML(false)
|
||||
enc.SetIndent(prefix, indent)
|
||||
if err := enc.Encode(&wire); err != nil {
|
||||
return nil, fmt.Errorf("marshaling jsonrpc message: %w", err)
|
||||
}
|
||||
return bytes.TrimRight(buf.Bytes(), "\n"), nil
|
||||
}
|
||||
|
||||
func DecodeMessage(data []byte) (Message, error) {
|
||||
msg := wireCombined{}
|
||||
if err := internaljson.Unmarshal(data, &msg); err != nil {
|
||||
return nil, fmt.Errorf("unmarshaling jsonrpc message: %w", err)
|
||||
}
|
||||
if msg.VersionTag != wireVersion {
|
||||
return nil, fmt.Errorf("invalid message version tag %q; expected %q", msg.VersionTag, wireVersion)
|
||||
}
|
||||
id, err := MakeID(msg.ID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if msg.Method != "" {
|
||||
// has a method, must be a call
|
||||
return &Request{
|
||||
Method: msg.Method,
|
||||
ID: id,
|
||||
Params: msg.Params,
|
||||
}, nil
|
||||
}
|
||||
// no method, should be a response
|
||||
if !id.IsValid() {
|
||||
return nil, ErrInvalidRequest
|
||||
}
|
||||
resp := &Response{
|
||||
ID: id,
|
||||
Result: msg.Result,
|
||||
}
|
||||
// we have to check if msg.Error is nil to avoid a typed error
|
||||
if msg.Error != nil {
|
||||
resp.Error = msg.Error
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
func marshalToRaw(obj any) (json.RawMessage, error) {
|
||||
if obj == nil {
|
||||
return nil, nil
|
||||
}
|
||||
data, err := jsonMarshal(obj)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return json.RawMessage(data), nil
|
||||
}
|
||||
|
||||
// jsonescaping is a compatibility parameter that allows to restore
|
||||
// JSON escaping in the JSON marshaling, which stopped being the default
|
||||
// in the 1.4.0 version of the SDK. See the documentation for the
|
||||
// mcpgodebug package for instructions how to enable it.
|
||||
// The option will be removed in the 1.6.0 version of the SDK.
|
||||
var jsonescaping = mcpgodebug.Value("jsonescaping")
|
||||
|
||||
// jsonMarshal marshals obj to JSON like json.Marshal but without HTML escaping.
|
||||
func jsonMarshal(obj any) ([]byte, error) {
|
||||
if jsonescaping == "1" {
|
||||
return json.Marshal(obj)
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
enc := json.NewEncoder(&buf)
|
||||
enc.SetEscapeHTML(false)
|
||||
if err := enc.Encode(obj); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// json.Encoder.Encode adds a trailing newline. Trim it to be consistent with json.Marshal.
|
||||
return bytes.TrimRight(buf.Bytes(), "\n"), nil
|
||||
}
|
||||
+138
@@ -0,0 +1,138 @@
|
||||
// Copyright 2018 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package jsonrpc2
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
)
|
||||
|
||||
// This file contains implementations of the transport primitives that use the standard network
|
||||
// package.
|
||||
|
||||
// NetListenOptions is the optional arguments to the NetListen function.
|
||||
type NetListenOptions struct {
|
||||
NetListenConfig net.ListenConfig
|
||||
NetDialer net.Dialer
|
||||
}
|
||||
|
||||
// NetListener returns a new Listener that listens on a socket using the net package.
|
||||
func NetListener(ctx context.Context, network, address string, options NetListenOptions) (Listener, error) {
|
||||
ln, err := options.NetListenConfig.Listen(ctx, network, address)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &netListener{net: ln}, nil
|
||||
}
|
||||
|
||||
// netListener is the implementation of Listener for connections made using the net package.
|
||||
type netListener struct {
|
||||
net net.Listener
|
||||
}
|
||||
|
||||
// Accept blocks waiting for an incoming connection to the listener.
|
||||
func (l *netListener) Accept(context.Context) (io.ReadWriteCloser, error) {
|
||||
return l.net.Accept()
|
||||
}
|
||||
|
||||
// Close will cause the listener to stop listening. It will not close any connections that have
|
||||
// already been accepted.
|
||||
func (l *netListener) Close() error {
|
||||
addr := l.net.Addr()
|
||||
err := l.net.Close()
|
||||
if addr.Network() == "unix" {
|
||||
rerr := os.Remove(addr.String())
|
||||
if rerr != nil && err == nil {
|
||||
err = rerr
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// Dialer returns a dialer that can be used to connect to the listener.
|
||||
func (l *netListener) Dialer() Dialer {
|
||||
return NetDialer(l.net.Addr().Network(), l.net.Addr().String(), net.Dialer{})
|
||||
}
|
||||
|
||||
// NetDialer returns a Dialer using the supplied standard network dialer.
|
||||
func NetDialer(network, address string, nd net.Dialer) Dialer {
|
||||
return &netDialer{
|
||||
network: network,
|
||||
address: address,
|
||||
dialer: nd,
|
||||
}
|
||||
}
|
||||
|
||||
type netDialer struct {
|
||||
network string
|
||||
address string
|
||||
dialer net.Dialer
|
||||
}
|
||||
|
||||
func (n *netDialer) Dial(ctx context.Context) (io.ReadWriteCloser, error) {
|
||||
return n.dialer.DialContext(ctx, n.network, n.address)
|
||||
}
|
||||
|
||||
// NetPipeListener returns a new Listener that listens using net.Pipe.
|
||||
// It is only possibly to connect to it using the Dialer returned by the
|
||||
// Dialer method, each call to that method will generate a new pipe the other
|
||||
// side of which will be returned from the Accept call.
|
||||
func NetPipeListener(ctx context.Context) (Listener, error) {
|
||||
return &netPiper{
|
||||
done: make(chan struct{}),
|
||||
dialed: make(chan io.ReadWriteCloser),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// netPiper is the implementation of Listener build on top of net.Pipes.
|
||||
type netPiper struct {
|
||||
done chan struct{}
|
||||
dialed chan io.ReadWriteCloser
|
||||
}
|
||||
|
||||
// Accept blocks waiting for an incoming connection to the listener.
|
||||
func (l *netPiper) Accept(context.Context) (io.ReadWriteCloser, error) {
|
||||
// Block until the pipe is dialed or the listener is closed,
|
||||
// preferring the latter if already closed at the start of Accept.
|
||||
select {
|
||||
case <-l.done:
|
||||
return nil, net.ErrClosed
|
||||
default:
|
||||
}
|
||||
select {
|
||||
case rwc := <-l.dialed:
|
||||
return rwc, nil
|
||||
case <-l.done:
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
}
|
||||
|
||||
// Close will cause the listener to stop listening. It will not close any connections that have
|
||||
// already been accepted.
|
||||
func (l *netPiper) Close() error {
|
||||
// unblock any accept calls that are pending
|
||||
close(l.done)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (l *netPiper) Dialer() Dialer {
|
||||
return l
|
||||
}
|
||||
|
||||
func (l *netPiper) Dial(ctx context.Context) (io.ReadWriteCloser, error) {
|
||||
client, server := net.Pipe()
|
||||
|
||||
select {
|
||||
case l.dialed <- server:
|
||||
return client, nil
|
||||
|
||||
case <-l.done:
|
||||
client.Close()
|
||||
server.Close()
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
}
|
||||
+330
@@ -0,0 +1,330 @@
|
||||
// Copyright 2020 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package jsonrpc2
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"runtime"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Listener is implemented by protocols to accept new inbound connections.
|
||||
type Listener interface {
|
||||
// Accept accepts an inbound connection to a server.
|
||||
// It blocks until either an inbound connection is made, or the listener is closed.
|
||||
Accept(context.Context) (io.ReadWriteCloser, error)
|
||||
|
||||
// Close closes the listener.
|
||||
// Any blocked Accept or Dial operations will unblock and return errors.
|
||||
Close() error
|
||||
|
||||
// Dialer returns a dialer that can be used to connect to this listener
|
||||
// locally.
|
||||
// If a listener does not implement this it will return nil.
|
||||
Dialer() Dialer
|
||||
}
|
||||
|
||||
// Dialer is used by clients to dial a server.
|
||||
type Dialer interface {
|
||||
// Dial returns a new communication byte stream to a listening server.
|
||||
Dial(ctx context.Context) (io.ReadWriteCloser, error)
|
||||
}
|
||||
|
||||
// Server is a running server that is accepting incoming connections.
|
||||
type Server struct {
|
||||
listener Listener
|
||||
binder Binder
|
||||
async *async
|
||||
|
||||
shutdownOnce sync.Once
|
||||
closing int32 // atomic: set to nonzero when Shutdown is called
|
||||
}
|
||||
|
||||
// Dial uses the dialer to make a new connection, wraps the returned
|
||||
// reader and writer using the framer to make a stream, and then builds
|
||||
// a connection on top of that stream using the binder.
|
||||
//
|
||||
// The returned Connection will operate independently using the Preempter and/or
|
||||
// Handler provided by the Binder, and will release its own resources when the
|
||||
// connection is broken, but the caller may Close it earlier to stop accepting
|
||||
// (or sending) new requests.
|
||||
//
|
||||
// If non-nil, the onDone function is called when the connection is closed.
|
||||
func Dial(ctx context.Context, dialer Dialer, binder Binder, onDone func()) (*Connection, error) {
|
||||
// dial a server
|
||||
rwc, err := dialer.Dial(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return bindConnection(ctx, rwc, binder, onDone), nil
|
||||
}
|
||||
|
||||
// NewServer starts a new server listening for incoming connections and returns
|
||||
// it.
|
||||
// This returns a fully running and connected server, it does not block on
|
||||
// the listener.
|
||||
// You can call Wait to block on the server, or Shutdown to get the sever to
|
||||
// terminate gracefully.
|
||||
// To notice incoming connections, use an intercepting Binder.
|
||||
func NewServer(ctx context.Context, listener Listener, binder Binder) *Server {
|
||||
server := &Server{
|
||||
listener: listener,
|
||||
binder: binder,
|
||||
async: newAsync(),
|
||||
}
|
||||
go server.run(ctx)
|
||||
return server
|
||||
}
|
||||
|
||||
// Wait returns only when the server has shut down.
|
||||
func (s *Server) Wait() error {
|
||||
return s.async.wait()
|
||||
}
|
||||
|
||||
// Shutdown informs the server to stop accepting new connections.
|
||||
func (s *Server) Shutdown() {
|
||||
s.shutdownOnce.Do(func() {
|
||||
atomic.StoreInt32(&s.closing, 1)
|
||||
s.listener.Close()
|
||||
})
|
||||
}
|
||||
|
||||
// run accepts incoming connections from the listener,
|
||||
// If IdleTimeout is non-zero, run exits after there are no clients for this
|
||||
// duration, otherwise it exits only on error.
|
||||
func (s *Server) run(ctx context.Context) {
|
||||
defer s.async.done()
|
||||
|
||||
var activeConns sync.WaitGroup
|
||||
for {
|
||||
rwc, err := s.listener.Accept(ctx)
|
||||
if err != nil {
|
||||
// Only Shutdown closes the listener. If we get an error after Shutdown is
|
||||
// called, assume that was the cause and don't report the error;
|
||||
// otherwise, report the error in case it is unexpected.
|
||||
if atomic.LoadInt32(&s.closing) == 0 {
|
||||
s.async.setError(err)
|
||||
}
|
||||
// We are done generating new connections for good.
|
||||
break
|
||||
}
|
||||
|
||||
// A new inbound connection.
|
||||
activeConns.Add(1)
|
||||
_ = bindConnection(ctx, rwc, s.binder, activeConns.Done) // unregisters itself when done
|
||||
}
|
||||
activeConns.Wait()
|
||||
}
|
||||
|
||||
// NewIdleListener wraps a listener with an idle timeout.
|
||||
//
|
||||
// When there are no active connections for at least the timeout duration,
|
||||
// calls to Accept will fail with ErrIdleTimeout.
|
||||
//
|
||||
// A connection is considered inactive as soon as its Close method is called.
|
||||
func NewIdleListener(timeout time.Duration, wrap Listener) Listener {
|
||||
l := &idleListener{
|
||||
wrapped: wrap,
|
||||
timeout: timeout,
|
||||
active: make(chan int, 1),
|
||||
timedOut: make(chan struct{}),
|
||||
idleTimer: make(chan *time.Timer, 1),
|
||||
}
|
||||
l.idleTimer <- time.AfterFunc(l.timeout, l.timerExpired)
|
||||
return l
|
||||
}
|
||||
|
||||
type idleListener struct {
|
||||
wrapped Listener
|
||||
timeout time.Duration
|
||||
|
||||
// Only one of these channels is receivable at any given time.
|
||||
active chan int // count of active connections; closed when Close is called if not timed out
|
||||
timedOut chan struct{} // closed when the idle timer expires
|
||||
idleTimer chan *time.Timer // holds the timer only when idle
|
||||
}
|
||||
|
||||
// Accept accepts an incoming connection.
|
||||
//
|
||||
// If an incoming connection is accepted concurrent to the listener being closed
|
||||
// due to idleness, the new connection is immediately closed.
|
||||
func (l *idleListener) Accept(ctx context.Context) (io.ReadWriteCloser, error) {
|
||||
rwc, err := l.wrapped.Accept(ctx)
|
||||
|
||||
select {
|
||||
case n, ok := <-l.active:
|
||||
if err != nil {
|
||||
if ok {
|
||||
l.active <- n
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if ok {
|
||||
l.active <- n + 1
|
||||
} else {
|
||||
// l.wrapped.Close Close has been called, but Accept returned a
|
||||
// connection. This race can occur with concurrent Accept and Close calls
|
||||
// with any net.Listener, and it is benign: since the listener was closed
|
||||
// explicitly, it can't have also timed out.
|
||||
}
|
||||
return l.newConn(rwc), nil
|
||||
|
||||
case <-l.timedOut:
|
||||
if err == nil {
|
||||
// Keeping the connection open would leave the listener simultaneously
|
||||
// active and closed due to idleness, which would be contradictory and
|
||||
// confusing. Close the connection and pretend that it never happened.
|
||||
rwc.Close()
|
||||
} else {
|
||||
// In theory the timeout could have raced with an unrelated error return
|
||||
// from Accept. However, ErrIdleTimeout is arguably still valid (since we
|
||||
// would have closed due to the timeout independent of the error), and the
|
||||
// harm from returning a spurious ErrIdleTimeout is negligible anyway.
|
||||
}
|
||||
return nil, ErrIdleTimeout
|
||||
|
||||
case timer := <-l.idleTimer:
|
||||
if err != nil {
|
||||
// The idle timer doesn't run until it receives itself from the idleTimer
|
||||
// channel, so it can't have called l.wrapped.Close yet and thus err can't
|
||||
// be ErrIdleTimeout. Leave the idle timer as it was and return whatever
|
||||
// error we got.
|
||||
l.idleTimer <- timer
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if !timer.Stop() {
|
||||
// Failed to stop the timer — the timer goroutine is in the process of
|
||||
// firing. Send the timer back to the timer goroutine so that it can
|
||||
// safely close the timedOut channel, and then wait for the listener to
|
||||
// actually be closed before we return ErrIdleTimeout.
|
||||
l.idleTimer <- timer
|
||||
rwc.Close()
|
||||
<-l.timedOut
|
||||
return nil, ErrIdleTimeout
|
||||
}
|
||||
|
||||
l.active <- 1
|
||||
return l.newConn(rwc), nil
|
||||
}
|
||||
}
|
||||
|
||||
func (l *idleListener) Close() error {
|
||||
select {
|
||||
case _, ok := <-l.active:
|
||||
if ok {
|
||||
close(l.active)
|
||||
}
|
||||
|
||||
case <-l.timedOut:
|
||||
// Already closed by the timer; take care not to double-close if the caller
|
||||
// only explicitly invokes this Close method once, since the io.Closer
|
||||
// interface explicitly leaves doubled Close calls undefined.
|
||||
return ErrIdleTimeout
|
||||
|
||||
case timer := <-l.idleTimer:
|
||||
if !timer.Stop() {
|
||||
// Couldn't stop the timer. It shouldn't take long to run, so just wait
|
||||
// (so that the Listener is guaranteed to be closed before we return)
|
||||
// and pretend that this call happened afterward.
|
||||
// That way we won't leak any timers or goroutines when Close returns.
|
||||
l.idleTimer <- timer
|
||||
<-l.timedOut
|
||||
return ErrIdleTimeout
|
||||
}
|
||||
close(l.active)
|
||||
}
|
||||
|
||||
return l.wrapped.Close()
|
||||
}
|
||||
|
||||
func (l *idleListener) Dialer() Dialer {
|
||||
return l.wrapped.Dialer()
|
||||
}
|
||||
|
||||
func (l *idleListener) timerExpired() {
|
||||
select {
|
||||
case n, ok := <-l.active:
|
||||
if ok {
|
||||
panic(fmt.Sprintf("jsonrpc2: idleListener idle timer fired with %d connections still active", n))
|
||||
} else {
|
||||
panic("jsonrpc2: Close finished with idle timer still running")
|
||||
}
|
||||
|
||||
case <-l.timedOut:
|
||||
panic("jsonrpc2: idleListener idle timer fired more than once")
|
||||
|
||||
case <-l.idleTimer:
|
||||
// The timer for this very call!
|
||||
}
|
||||
|
||||
// Close the Listener with all channels still blocked to ensure that this call
|
||||
// to l.wrapped.Close doesn't race with the one in l.Close.
|
||||
defer close(l.timedOut)
|
||||
l.wrapped.Close()
|
||||
}
|
||||
|
||||
func (l *idleListener) connClosed() {
|
||||
select {
|
||||
case n, ok := <-l.active:
|
||||
if !ok {
|
||||
// l is already closed, so it can't close due to idleness,
|
||||
// and we don't need to track the number of active connections any more.
|
||||
return
|
||||
}
|
||||
n--
|
||||
if n == 0 {
|
||||
l.idleTimer <- time.AfterFunc(l.timeout, l.timerExpired)
|
||||
} else {
|
||||
l.active <- n
|
||||
}
|
||||
|
||||
case <-l.timedOut:
|
||||
panic("jsonrpc2: idleListener idle timer fired before last active connection was closed")
|
||||
|
||||
case <-l.idleTimer:
|
||||
panic("jsonrpc2: idleListener idle timer active before last active connection was closed")
|
||||
}
|
||||
}
|
||||
|
||||
type idleListenerConn struct {
|
||||
wrapped io.ReadWriteCloser
|
||||
l *idleListener
|
||||
closeOnce sync.Once
|
||||
}
|
||||
|
||||
func (l *idleListener) newConn(rwc io.ReadWriteCloser) *idleListenerConn {
|
||||
c := &idleListenerConn{
|
||||
wrapped: rwc,
|
||||
l: l,
|
||||
}
|
||||
|
||||
// A caller that forgets to call Close may disrupt the idleListener's
|
||||
// accounting, even though the file descriptor for the underlying connection
|
||||
// may eventually be garbage-collected anyway.
|
||||
//
|
||||
// Set a (best-effort) finalizer to verify that a Close call always occurs.
|
||||
// (We will clear the finalizer explicitly in Close.)
|
||||
runtime.SetFinalizer(c, func(c *idleListenerConn) {
|
||||
panic("jsonrpc2: IdleListener connection became unreachable without a call to Close")
|
||||
})
|
||||
|
||||
return c
|
||||
}
|
||||
|
||||
func (c *idleListenerConn) Read(p []byte) (int, error) { return c.wrapped.Read(p) }
|
||||
func (c *idleListenerConn) Write(p []byte) (int, error) { return c.wrapped.Write(p) }
|
||||
|
||||
func (c *idleListenerConn) Close() error {
|
||||
defer c.closeOnce.Do(func() {
|
||||
c.l.connClosed()
|
||||
runtime.SetFinalizer(c, nil)
|
||||
})
|
||||
return c.wrapped.Close()
|
||||
}
|
||||
+97
@@ -0,0 +1,97 @@
|
||||
// Copyright 2018 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package jsonrpc2
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
)
|
||||
|
||||
// This file contains the go forms of the wire specification.
|
||||
// see http://www.jsonrpc.org/specification for details
|
||||
|
||||
var (
|
||||
// ErrParse is used when invalid JSON was received by the server.
|
||||
ErrParse = NewError(-32700, "parse error")
|
||||
// ErrInvalidRequest is used when the JSON sent is not a valid Request object.
|
||||
ErrInvalidRequest = NewError(-32600, "invalid request")
|
||||
// ErrMethodNotFound should be returned by the handler when the method does
|
||||
// not exist / is not available.
|
||||
ErrMethodNotFound = NewError(-32601, "method not found")
|
||||
// ErrInvalidParams should be returned by the handler when method
|
||||
// parameter(s) were invalid.
|
||||
ErrInvalidParams = NewError(-32602, "invalid params")
|
||||
// ErrInternal indicates a failure to process a call correctly
|
||||
ErrInternal = NewError(-32603, "internal error")
|
||||
|
||||
// The following errors are not part of the json specification, but
|
||||
// compliant extensions specific to this implementation.
|
||||
|
||||
// ErrServerOverloaded is returned when a message was refused due to a
|
||||
// server being temporarily unable to accept any new messages.
|
||||
ErrServerOverloaded = NewError(-32000, "overloaded")
|
||||
// ErrUnknown should be used for all non coded errors.
|
||||
ErrUnknown = NewError(-32001, "unknown error")
|
||||
// ErrServerClosing is returned for calls that arrive while the server is closing.
|
||||
ErrServerClosing = NewError(-32004, "server is closing")
|
||||
// ErrClientClosing is a dummy error returned for calls initiated while the client is closing.
|
||||
ErrClientClosing = NewError(-32003, "client is closing")
|
||||
|
||||
// The following errors have special semantics for MCP transports
|
||||
|
||||
// ErrRejected may be wrapped to return errors from calls to Writer.Write
|
||||
// that signal that the request was rejected by the transport layer as
|
||||
// invalid.
|
||||
//
|
||||
// Such failures do not indicate that the connection is broken, but rather
|
||||
// should be returned to the caller to indicate that the specific request is
|
||||
// invalid in the current context.
|
||||
ErrRejected = NewError(-32005, "rejected by transport")
|
||||
)
|
||||
|
||||
const wireVersion = "2.0"
|
||||
|
||||
// wireCombined has all the fields of both Request and Response.
|
||||
// We can decode this and then work out which it is.
|
||||
type wireCombined struct {
|
||||
VersionTag string `json:"jsonrpc"`
|
||||
ID any `json:"id,omitempty"`
|
||||
Method string `json:"method,omitempty"`
|
||||
Params json.RawMessage `json:"params,omitempty"`
|
||||
Result json.RawMessage `json:"result,omitempty"`
|
||||
Error *WireError `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
// WireError represents a structured error in a Response.
|
||||
type WireError struct {
|
||||
// Code is an error code indicating the type of failure.
|
||||
Code int64 `json:"code"`
|
||||
// Message is a short description of the error.
|
||||
Message string `json:"message"`
|
||||
// Data is optional structured data containing additional information about the error.
|
||||
Data json.RawMessage `json:"data,omitempty"`
|
||||
}
|
||||
|
||||
// NewError returns an error that will encode on the wire correctly.
|
||||
// The standard codes are made available from this package, this function should
|
||||
// only be used to build errors for application specific codes as allowed by the
|
||||
// specification.
|
||||
func NewError(code int64, message string) error {
|
||||
return &WireError{
|
||||
Code: code,
|
||||
Message: message,
|
||||
}
|
||||
}
|
||||
|
||||
func (err *WireError) Error() string {
|
||||
return err.Message
|
||||
}
|
||||
|
||||
func (err *WireError) Is(other error) bool {
|
||||
w, ok := other.(*WireError)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return err.Code == w.Code
|
||||
}
|
||||
+52
@@ -0,0 +1,52 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by the license
|
||||
// that can be found in the LICENSE file.
|
||||
|
||||
// Package mcpgodebug provides a mechanism to configure compatibility parameters
|
||||
// via the MCPGODEBUG environment variable.
|
||||
//
|
||||
// The value of MCPGODEBUG is a comma-separated list of key=value pairs.
|
||||
// For example:
|
||||
//
|
||||
// MCPGODEBUG=someoption=1,otheroption=value
|
||||
package mcpgodebug
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const compatibilityEnvKey = "MCPGODEBUG"
|
||||
|
||||
var compatibilityParams map[string]string
|
||||
|
||||
func init() {
|
||||
var err error
|
||||
compatibilityParams, err = parseCompatibility(os.Getenv(compatibilityEnvKey))
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
|
||||
// Value returns the value of the compatibility parameter with the given key.
|
||||
// It returns an empty string if the key is not set.
|
||||
func Value(key string) string {
|
||||
return compatibilityParams[key]
|
||||
}
|
||||
|
||||
func parseCompatibility(envValue string) (map[string]string, error) {
|
||||
if envValue == "" {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
params := make(map[string]string)
|
||||
for part := range strings.SplitSeq(envValue, ",") {
|
||||
k, v, ok := strings.Cut(part, "=")
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("MCPGODEBUG: invalid format: %q", part)
|
||||
}
|
||||
params[strings.TrimSpace(k)] = strings.TrimSpace(v)
|
||||
}
|
||||
return params, nil
|
||||
}
|
||||
+26
@@ -0,0 +1,26 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by the license
|
||||
// that can be found in the LICENSE file.
|
||||
package util
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func IsLoopback(addr string) bool {
|
||||
host, _, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
// If SplitHostPort fails, it might be just a host without a port.
|
||||
host = strings.Trim(addr, "[]")
|
||||
}
|
||||
if host == "localhost" {
|
||||
return true
|
||||
}
|
||||
ip, err := netip.ParseAddr(host)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return ip.IsLoopback()
|
||||
}
|
||||
+44
@@ -0,0 +1,44 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"cmp"
|
||||
"fmt"
|
||||
"iter"
|
||||
"slices"
|
||||
)
|
||||
|
||||
// Helpers below are copied from gopls' moremaps package.
|
||||
|
||||
// Sorted returns an iterator over the entries of m in key order.
|
||||
func Sorted[M ~map[K]V, K cmp.Ordered, V any](m M) iter.Seq2[K, V] {
|
||||
// TODO(adonovan): use maps.Sorted if proposal #68598 is accepted.
|
||||
return func(yield func(K, V) bool) {
|
||||
keys := KeySlice(m)
|
||||
slices.Sort(keys)
|
||||
for _, k := range keys {
|
||||
if !yield(k, m[k]) {
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// KeySlice returns the keys of the map M, like slices.Collect(maps.Keys(m)).
|
||||
func KeySlice[M ~map[K]V, K comparable, V any](m M) []K {
|
||||
r := make([]K, 0, len(m))
|
||||
for k := range m {
|
||||
r = append(r, k)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// Wrapf wraps *errp with the given formatted message if *errp is not nil.
|
||||
func Wrapf(errp *error, format string, args ...any) {
|
||||
if *errp != nil {
|
||||
*errp = fmt.Errorf("%s: %w", fmt.Sprintf(format, args...), *errp)
|
||||
}
|
||||
}
|
||||
+23
@@ -0,0 +1,23 @@
|
||||
// Copyright 2019 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// Package xcontext is a package to offer the extra functionality we need
|
||||
// from contexts that is not available from the standard context package.
|
||||
package xcontext
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Detach returns a context that keeps all the values of its parent context
|
||||
// but detaches from the cancellation and error handling.
|
||||
func Detach(ctx context.Context) context.Context { return detachedContext{ctx} }
|
||||
|
||||
type detachedContext struct{ parent context.Context }
|
||||
|
||||
func (v detachedContext) Deadline() (time.Time, bool) { return time.Time{}, false }
|
||||
func (v detachedContext) Done() <-chan struct{} { return nil }
|
||||
func (v detachedContext) Err() error { return nil }
|
||||
func (v detachedContext) Value(key any) any { return v.parent.Value(key) }
|
||||
+56
@@ -0,0 +1,56 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// Package jsonrpc exposes part of a JSON-RPC v2 implementation
|
||||
// for use by mcp transport authors.
|
||||
package jsonrpc
|
||||
|
||||
import "github.com/modelcontextprotocol/go-sdk/internal/jsonrpc2"
|
||||
|
||||
type (
|
||||
// ID is a JSON-RPC request ID.
|
||||
ID = jsonrpc2.ID
|
||||
// Message is a JSON-RPC message.
|
||||
Message = jsonrpc2.Message
|
||||
// Request is a JSON-RPC request.
|
||||
Request = jsonrpc2.Request
|
||||
// Response is a JSON-RPC response.
|
||||
Response = jsonrpc2.Response
|
||||
// Error is a structured error in a JSON-RPC response.
|
||||
Error = jsonrpc2.WireError
|
||||
)
|
||||
|
||||
// MakeID coerces the given Go value to an ID. The value should be the
|
||||
// default JSON marshaling of a Request identifier: nil, float64, or string.
|
||||
//
|
||||
// Returns an error if the value type was not a valid Request ID type.
|
||||
func MakeID(v any) (ID, error) {
|
||||
return jsonrpc2.MakeID(v)
|
||||
}
|
||||
|
||||
// EncodeMessage serializes a JSON-RPC message to its wire format.
|
||||
func EncodeMessage(msg Message) ([]byte, error) {
|
||||
return jsonrpc2.EncodeMessage(msg)
|
||||
}
|
||||
|
||||
// DecodeMessage deserializes JSON-RPC wire format data into a Message.
|
||||
// It returns either a Request or Response based on the message content.
|
||||
func DecodeMessage(data []byte) (Message, error) {
|
||||
return jsonrpc2.DecodeMessage(data)
|
||||
}
|
||||
|
||||
// Standard JSON-RPC 2.0 error codes.
|
||||
// See https://www.jsonrpc.org/specification#error_object
|
||||
const (
|
||||
// CodeParseError indicates invalid JSON was received by the server.
|
||||
CodeParseError = -32700
|
||||
// CodeInvalidRequest indicates the JSON sent is not a valid Request object.
|
||||
CodeInvalidRequest = -32600
|
||||
// CodeMethodNotFound indicates the method does not exist or is not available.
|
||||
CodeMethodNotFound = -32601
|
||||
// CodeInvalidParams indicates invalid method parameter(s).
|
||||
CodeInvalidParams = -32602
|
||||
// CodeInternalError indicates an internal JSON-RPC error.
|
||||
CodeInternalError = -32603
|
||||
)
|
||||
+1182
File diff suppressed because it is too large
Load Diff
+108
@@ -0,0 +1,108 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"os/exec"
|
||||
"syscall"
|
||||
"time"
|
||||
)
|
||||
|
||||
var defaultTerminateDuration = 5 * time.Second // mutable for testing
|
||||
|
||||
// A CommandTransport is a [Transport] that runs a command and communicates
|
||||
// with it over stdin/stdout, using newline-delimited JSON.
|
||||
type CommandTransport struct {
|
||||
Command *exec.Cmd
|
||||
// TerminateDuration controls how long Close waits after closing stdin
|
||||
// for the process to exit before sending SIGTERM.
|
||||
// If zero or negative, the default of 5s is used.
|
||||
TerminateDuration time.Duration
|
||||
}
|
||||
|
||||
// Connect starts the command, and connects to it over stdin/stdout.
|
||||
func (t *CommandTransport) Connect(ctx context.Context) (Connection, error) {
|
||||
stdout, err := t.Command.StdoutPipe()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stdout = io.NopCloser(stdout) // close the connection by closing stdin, not stdout
|
||||
stdin, err := t.Command.StdinPipe()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := t.Command.Start(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
td := t.TerminateDuration
|
||||
if td <= 0 {
|
||||
td = defaultTerminateDuration
|
||||
}
|
||||
return newIOConn(&pipeRWC{t.Command, stdout, stdin, td}), nil
|
||||
}
|
||||
|
||||
// A pipeRWC is an io.ReadWriteCloser that communicates with a subprocess over
|
||||
// stdin/stdout pipes.
|
||||
type pipeRWC struct {
|
||||
cmd *exec.Cmd
|
||||
stdout io.ReadCloser
|
||||
stdin io.WriteCloser
|
||||
terminateDuration time.Duration
|
||||
}
|
||||
|
||||
func (s *pipeRWC) Read(p []byte) (n int, err error) {
|
||||
return s.stdout.Read(p)
|
||||
}
|
||||
|
||||
func (s *pipeRWC) Write(p []byte) (n int, err error) {
|
||||
return s.stdin.Write(p)
|
||||
}
|
||||
|
||||
// Close closes the input stream to the child process, and awaits normal
|
||||
// termination of the command. If the command does not exit, it is signalled to
|
||||
// terminate, and then eventually killed.
|
||||
func (s *pipeRWC) Close() error {
|
||||
// Spec:
|
||||
// "For the stdio transport, the client SHOULD initiate shutdown by:...
|
||||
|
||||
// "...First, closing the input stream to the child process (the server)"
|
||||
if err := s.stdin.Close(); err != nil {
|
||||
return fmt.Errorf("closing stdin: %v", err)
|
||||
}
|
||||
resChan := make(chan error, 1)
|
||||
go func() {
|
||||
resChan <- s.cmd.Wait()
|
||||
}()
|
||||
// "...Waiting for the server to exit, or sending SIGTERM if the server does not exit within a reasonable time"
|
||||
wait := func() (error, bool) {
|
||||
select {
|
||||
case err := <-resChan:
|
||||
return err, true
|
||||
case <-time.After(s.terminateDuration):
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
if err, ok := wait(); ok {
|
||||
return err
|
||||
}
|
||||
// Note the condition here: if sending SIGTERM fails, don't wait and just
|
||||
// move on to SIGKILL.
|
||||
if err := s.cmd.Process.Signal(syscall.SIGTERM); err == nil {
|
||||
if err, ok := wait(); ok {
|
||||
return err
|
||||
}
|
||||
}
|
||||
// "...Sending SIGKILL if the server does not exit within a reasonable time after SIGTERM"
|
||||
if err := s.cmd.Process.Kill(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err, ok := wait(); ok {
|
||||
return err
|
||||
}
|
||||
return fmt.Errorf("unresponsive subprocess")
|
||||
}
|
||||
+410
@@ -0,0 +1,410 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// TODO(findleyr): update JSON marshalling of all content types to preserve required fields.
|
||||
// (See [TextContent.MarshalJSON], which handles this for text content).
|
||||
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
internaljson "github.com/modelcontextprotocol/go-sdk/internal/json"
|
||||
)
|
||||
|
||||
// A Content is a [TextContent], [ImageContent], [AudioContent],
|
||||
// [ResourceLink], [EmbeddedResource], [ToolUseContent], or [ToolResultContent].
|
||||
//
|
||||
// Note: [ToolUseContent] and [ToolResultContent] are only valid in sampling
|
||||
// message contexts (CreateMessageParams/CreateMessageResult).
|
||||
type Content interface {
|
||||
MarshalJSON() ([]byte, error)
|
||||
fromWire(*wireContent)
|
||||
}
|
||||
|
||||
// TextContent is a textual content.
|
||||
type TextContent struct {
|
||||
Text string
|
||||
Meta Meta
|
||||
Annotations *Annotations
|
||||
}
|
||||
|
||||
func (c *TextContent) MarshalJSON() ([]byte, error) {
|
||||
// Custom wire format to ensure the required "text" field is always included, even when empty.
|
||||
wire := struct {
|
||||
Type string `json:"type"`
|
||||
Text string `json:"text"`
|
||||
Meta Meta `json:"_meta,omitempty"`
|
||||
Annotations *Annotations `json:"annotations,omitempty"`
|
||||
}{
|
||||
Type: "text",
|
||||
Text: c.Text,
|
||||
Meta: c.Meta,
|
||||
Annotations: c.Annotations,
|
||||
}
|
||||
return json.Marshal(wire)
|
||||
}
|
||||
|
||||
func (c *TextContent) fromWire(wire *wireContent) {
|
||||
c.Text = wire.Text
|
||||
c.Meta = wire.Meta
|
||||
c.Annotations = wire.Annotations
|
||||
}
|
||||
|
||||
// ImageContent contains base64-encoded image data.
|
||||
type ImageContent struct {
|
||||
Meta Meta
|
||||
Annotations *Annotations
|
||||
Data []byte // base64-encoded
|
||||
MIMEType string
|
||||
}
|
||||
|
||||
func (c *ImageContent) MarshalJSON() ([]byte, error) {
|
||||
// Custom wire format to ensure required fields are always included, even when empty.
|
||||
data := c.Data
|
||||
if data == nil {
|
||||
data = []byte{}
|
||||
}
|
||||
wire := imageAudioWire{
|
||||
Type: "image",
|
||||
MIMEType: c.MIMEType,
|
||||
Data: data,
|
||||
Meta: c.Meta,
|
||||
Annotations: c.Annotations,
|
||||
}
|
||||
return json.Marshal(wire)
|
||||
}
|
||||
|
||||
func (c *ImageContent) fromWire(wire *wireContent) {
|
||||
c.MIMEType = wire.MIMEType
|
||||
c.Data = wire.Data
|
||||
c.Meta = wire.Meta
|
||||
c.Annotations = wire.Annotations
|
||||
}
|
||||
|
||||
// AudioContent contains base64-encoded audio data.
|
||||
type AudioContent struct {
|
||||
Data []byte
|
||||
MIMEType string
|
||||
Meta Meta
|
||||
Annotations *Annotations
|
||||
}
|
||||
|
||||
func (c AudioContent) MarshalJSON() ([]byte, error) {
|
||||
// Custom wire format to ensure required fields are always included, even when empty.
|
||||
data := c.Data
|
||||
if data == nil {
|
||||
data = []byte{}
|
||||
}
|
||||
wire := imageAudioWire{
|
||||
Type: "audio",
|
||||
MIMEType: c.MIMEType,
|
||||
Data: data,
|
||||
Meta: c.Meta,
|
||||
Annotations: c.Annotations,
|
||||
}
|
||||
return json.Marshal(wire)
|
||||
}
|
||||
|
||||
func (c *AudioContent) fromWire(wire *wireContent) {
|
||||
c.MIMEType = wire.MIMEType
|
||||
c.Data = wire.Data
|
||||
c.Meta = wire.Meta
|
||||
c.Annotations = wire.Annotations
|
||||
}
|
||||
|
||||
// Custom wire format to ensure required fields are always included, even when empty.
|
||||
type imageAudioWire struct {
|
||||
Type string `json:"type"`
|
||||
MIMEType string `json:"mimeType"`
|
||||
Data []byte `json:"data"`
|
||||
Meta Meta `json:"_meta,omitempty"`
|
||||
Annotations *Annotations `json:"annotations,omitempty"`
|
||||
}
|
||||
|
||||
// ResourceLink is a link to a resource
|
||||
type ResourceLink struct {
|
||||
URI string
|
||||
Name string
|
||||
Title string
|
||||
Description string
|
||||
MIMEType string
|
||||
Size *int64
|
||||
Meta Meta
|
||||
Annotations *Annotations
|
||||
// Icons for the resource link, if any.
|
||||
Icons []Icon `json:"icons,omitempty"`
|
||||
}
|
||||
|
||||
func (c *ResourceLink) MarshalJSON() ([]byte, error) {
|
||||
return json.Marshal(&wireContent{
|
||||
Type: "resource_link",
|
||||
URI: c.URI,
|
||||
Name: c.Name,
|
||||
Title: c.Title,
|
||||
Description: c.Description,
|
||||
MIMEType: c.MIMEType,
|
||||
Size: c.Size,
|
||||
Meta: c.Meta,
|
||||
Annotations: c.Annotations,
|
||||
Icons: c.Icons,
|
||||
})
|
||||
}
|
||||
|
||||
func (c *ResourceLink) fromWire(wire *wireContent) {
|
||||
c.URI = wire.URI
|
||||
c.Name = wire.Name
|
||||
c.Title = wire.Title
|
||||
c.Description = wire.Description
|
||||
c.MIMEType = wire.MIMEType
|
||||
c.Size = wire.Size
|
||||
c.Meta = wire.Meta
|
||||
c.Annotations = wire.Annotations
|
||||
c.Icons = wire.Icons
|
||||
}
|
||||
|
||||
// EmbeddedResource contains embedded resources.
|
||||
type EmbeddedResource struct {
|
||||
Resource *ResourceContents
|
||||
Meta Meta
|
||||
Annotations *Annotations
|
||||
}
|
||||
|
||||
func (c *EmbeddedResource) MarshalJSON() ([]byte, error) {
|
||||
return json.Marshal(&wireContent{
|
||||
Type: "resource",
|
||||
Resource: c.Resource,
|
||||
Meta: c.Meta,
|
||||
Annotations: c.Annotations,
|
||||
})
|
||||
}
|
||||
|
||||
func (c *EmbeddedResource) fromWire(wire *wireContent) {
|
||||
c.Resource = wire.Resource
|
||||
c.Meta = wire.Meta
|
||||
c.Annotations = wire.Annotations
|
||||
}
|
||||
|
||||
// ToolUseContent represents a request from the assistant to invoke a tool.
|
||||
// This content type is only valid in sampling messages.
|
||||
type ToolUseContent struct {
|
||||
// ID is a unique identifier for this tool use, used to match with ToolResultContent.
|
||||
ID string
|
||||
// Name is the name of the tool to invoke.
|
||||
Name string
|
||||
// Input contains the tool arguments as a JSON object.
|
||||
Input map[string]any
|
||||
Meta Meta
|
||||
}
|
||||
|
||||
func (c *ToolUseContent) MarshalJSON() ([]byte, error) {
|
||||
input := c.Input
|
||||
if input == nil {
|
||||
input = map[string]any{}
|
||||
}
|
||||
wire := struct {
|
||||
Type string `json:"type"`
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Input map[string]any `json:"input"`
|
||||
Meta Meta `json:"_meta,omitempty"`
|
||||
}{
|
||||
Type: "tool_use",
|
||||
ID: c.ID,
|
||||
Name: c.Name,
|
||||
Input: input,
|
||||
Meta: c.Meta,
|
||||
}
|
||||
return json.Marshal(wire)
|
||||
}
|
||||
|
||||
func (c *ToolUseContent) fromWire(wire *wireContent) {
|
||||
c.ID = wire.ID
|
||||
c.Name = wire.Name
|
||||
c.Input = wire.Input
|
||||
c.Meta = wire.Meta
|
||||
}
|
||||
|
||||
// ToolResultContent represents the result of a tool invocation.
|
||||
// This content type is only valid in sampling messages with role "user".
|
||||
type ToolResultContent struct {
|
||||
// ToolUseID references the ID from the corresponding ToolUseContent.
|
||||
ToolUseID string
|
||||
// Content holds the unstructured result of the tool call.
|
||||
Content []Content
|
||||
// StructuredContent holds an optional structured result as a JSON object.
|
||||
StructuredContent any
|
||||
// IsError indicates whether the tool call ended in an error.
|
||||
IsError bool
|
||||
Meta Meta
|
||||
}
|
||||
|
||||
func (c *ToolResultContent) MarshalJSON() ([]byte, error) {
|
||||
// Marshal nested content
|
||||
var contentWire []*wireContent
|
||||
for _, content := range c.Content {
|
||||
data, err := content.MarshalJSON()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var w wireContent
|
||||
if err := internaljson.Unmarshal(data, &w); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
contentWire = append(contentWire, &w)
|
||||
}
|
||||
if contentWire == nil {
|
||||
contentWire = []*wireContent{} // avoid JSON null
|
||||
}
|
||||
|
||||
wire := struct {
|
||||
Type string `json:"type"`
|
||||
ToolUseID string `json:"toolUseId"`
|
||||
Content []*wireContent `json:"content"`
|
||||
StructuredContent any `json:"structuredContent,omitempty"`
|
||||
IsError bool `json:"isError,omitempty"`
|
||||
Meta Meta `json:"_meta,omitempty"`
|
||||
}{
|
||||
Type: "tool_result",
|
||||
ToolUseID: c.ToolUseID,
|
||||
Content: contentWire,
|
||||
StructuredContent: c.StructuredContent,
|
||||
IsError: c.IsError,
|
||||
Meta: c.Meta,
|
||||
}
|
||||
return json.Marshal(wire)
|
||||
}
|
||||
|
||||
func (c *ToolResultContent) fromWire(wire *wireContent) {
|
||||
c.ToolUseID = wire.ToolUseID
|
||||
c.StructuredContent = wire.StructuredContent
|
||||
c.IsError = wire.IsError
|
||||
c.Meta = wire.Meta
|
||||
// Content is handled separately in contentFromWire due to nested content
|
||||
}
|
||||
|
||||
// ResourceContents contains the contents of a specific resource or
|
||||
// sub-resource.
|
||||
type ResourceContents struct {
|
||||
URI string `json:"uri"`
|
||||
MIMEType string `json:"mimeType,omitempty"`
|
||||
Text string `json:"text,omitempty"`
|
||||
Blob []byte `json:"blob,omitzero"`
|
||||
Meta Meta `json:"_meta,omitempty"`
|
||||
}
|
||||
|
||||
// wireContent is the wire format for content.
|
||||
// It represents the protocol types TextContent, ImageContent, AudioContent,
|
||||
// ResourceLink, EmbeddedResource, ToolUseContent, and ToolResultContent.
|
||||
// The Type field distinguishes them. In the protocol, each type has a constant
|
||||
// value for the field.
|
||||
type wireContent struct {
|
||||
Type string `json:"type"`
|
||||
Text string `json:"text,omitempty"` // TextContent
|
||||
MIMEType string `json:"mimeType,omitempty"` // ImageContent, AudioContent, ResourceLink
|
||||
Data []byte `json:"data,omitempty"` // ImageContent, AudioContent
|
||||
Resource *ResourceContents `json:"resource,omitempty"` // EmbeddedResource
|
||||
URI string `json:"uri,omitempty"` // ResourceLink
|
||||
Name string `json:"name,omitempty"` // ResourceLink, ToolUseContent
|
||||
Title string `json:"title,omitempty"` // ResourceLink
|
||||
Description string `json:"description,omitempty"` // ResourceLink
|
||||
Size *int64 `json:"size,omitempty"` // ResourceLink
|
||||
Meta Meta `json:"_meta,omitempty"` // all types
|
||||
Annotations *Annotations `json:"annotations,omitempty"` // all types except ToolUseContent, ToolResultContent
|
||||
Icons []Icon `json:"icons,omitempty"` // ResourceLink
|
||||
ID string `json:"id,omitempty"` // ToolUseContent
|
||||
Input map[string]any `json:"input,omitempty"` // ToolUseContent
|
||||
ToolUseID string `json:"toolUseId,omitempty"` // ToolResultContent
|
||||
NestedContent []*wireContent `json:"content,omitempty"` // ToolResultContent
|
||||
StructuredContent any `json:"structuredContent,omitempty"` // ToolResultContent
|
||||
IsError bool `json:"isError,omitempty"` // ToolResultContent
|
||||
}
|
||||
|
||||
// unmarshalContent unmarshals JSON that is either a single content object or
|
||||
// an array of content objects. A single object is wrapped in a one-element slice.
|
||||
func unmarshalContent(raw json.RawMessage, allow map[string]bool) ([]Content, error) {
|
||||
if len(raw) == 0 || string(raw) == "null" {
|
||||
return nil, fmt.Errorf("nil content")
|
||||
}
|
||||
// Try array first, then fall back to single object.
|
||||
var wires []*wireContent
|
||||
if err := internaljson.Unmarshal(raw, &wires); err == nil {
|
||||
return contentsFromWire(wires, allow)
|
||||
}
|
||||
var wire wireContent
|
||||
if err := internaljson.Unmarshal(raw, &wire); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c, err := contentFromWire(&wire, allow)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []Content{c}, nil
|
||||
}
|
||||
|
||||
func contentsFromWire(wires []*wireContent, allow map[string]bool) ([]Content, error) {
|
||||
blocks := make([]Content, 0, len(wires))
|
||||
for _, wire := range wires {
|
||||
block, err := contentFromWire(wire, allow)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
blocks = append(blocks, block)
|
||||
}
|
||||
return blocks, nil
|
||||
}
|
||||
|
||||
func contentFromWire(wire *wireContent, allow map[string]bool) (Content, error) {
|
||||
if wire == nil {
|
||||
return nil, fmt.Errorf("nil content")
|
||||
}
|
||||
if allow != nil && !allow[wire.Type] {
|
||||
return nil, fmt.Errorf("invalid content type %q", wire.Type)
|
||||
}
|
||||
switch wire.Type {
|
||||
case "text":
|
||||
v := new(TextContent)
|
||||
v.fromWire(wire)
|
||||
return v, nil
|
||||
case "image":
|
||||
v := new(ImageContent)
|
||||
v.fromWire(wire)
|
||||
return v, nil
|
||||
case "audio":
|
||||
v := new(AudioContent)
|
||||
v.fromWire(wire)
|
||||
return v, nil
|
||||
case "resource_link":
|
||||
v := new(ResourceLink)
|
||||
v.fromWire(wire)
|
||||
return v, nil
|
||||
case "resource":
|
||||
v := new(EmbeddedResource)
|
||||
v.fromWire(wire)
|
||||
return v, nil
|
||||
case "tool_use":
|
||||
v := new(ToolUseContent)
|
||||
v.fromWire(wire)
|
||||
return v, nil
|
||||
case "tool_result":
|
||||
v := new(ToolResultContent)
|
||||
v.fromWire(wire)
|
||||
// Handle nested content - tool_result content can contain text, image, audio,
|
||||
// resource_link, and resource (same as CallToolResult.content)
|
||||
if wire.NestedContent != nil {
|
||||
toolResultContentAllow := map[string]bool{
|
||||
"text": true, "image": true, "audio": true,
|
||||
"resource_link": true, "resource": true,
|
||||
}
|
||||
nestedContent, err := contentsFromWire(wire.NestedContent, toolResultContentAllow)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("tool_result nested content: %w", err)
|
||||
}
|
||||
v.Content = nestedContent
|
||||
}
|
||||
return v, nil
|
||||
}
|
||||
return nil, fmt.Errorf("unrecognized content type %q", wire.Type)
|
||||
}
|
||||
+436
@@ -0,0 +1,436 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// This file is for SSE events.
|
||||
// See https://developer.mozilla.org/en-US/docs/Web/API/Server-sent_events/Using_server-sent_events.
|
||||
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"iter"
|
||||
"maps"
|
||||
"net/http"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// If true, MemoryEventStore will do frequent validation to check invariants, slowing it down.
|
||||
// Enable for debugging.
|
||||
const validateMemoryEventStore = false
|
||||
|
||||
// An Event is a server-sent event.
|
||||
// See https://developer.mozilla.org/en-US/docs/Web/API/Server-sent_events/Using_server-sent_events#fields.
|
||||
type Event struct {
|
||||
Name string // the "event" field
|
||||
ID string // the "id" field
|
||||
Data []byte // the "data" field
|
||||
Retry string // the "retry" field
|
||||
}
|
||||
|
||||
// Empty reports whether the Event is empty.
|
||||
func (e Event) Empty() bool {
|
||||
return e.Name == "" && e.ID == "" && len(e.Data) == 0 && e.Retry == ""
|
||||
}
|
||||
|
||||
// writeEvent writes the event to w, and flushes.
|
||||
func writeEvent(w io.Writer, evt Event) (int, error) {
|
||||
var b bytes.Buffer
|
||||
if evt.Name != "" {
|
||||
fmt.Fprintf(&b, "event: %s\n", evt.Name)
|
||||
}
|
||||
if evt.ID != "" {
|
||||
fmt.Fprintf(&b, "id: %s\n", evt.ID)
|
||||
}
|
||||
if evt.Retry != "" {
|
||||
fmt.Fprintf(&b, "retry: %s\n", evt.Retry)
|
||||
}
|
||||
fmt.Fprintf(&b, "data: %s\n\n", string(evt.Data))
|
||||
n, err := w.Write(b.Bytes())
|
||||
if f, ok := w.(http.Flusher); ok {
|
||||
f.Flush()
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
// scanEvents iterates SSE events in the given scanner. The iterated error is
|
||||
// terminal: if encountered, the stream is corrupt or broken and should no
|
||||
// longer be used.
|
||||
//
|
||||
// TODO(rfindley): consider a different API here that makes failure modes more
|
||||
// apparent.
|
||||
func scanEvents(r io.Reader) iter.Seq2[Event, error] {
|
||||
reader := bufio.NewReader(r)
|
||||
|
||||
// TODO: investigate proper behavior when events are out of order, or have
|
||||
// non-standard names.
|
||||
var (
|
||||
eventKey = []byte("event")
|
||||
idKey = []byte("id")
|
||||
dataKey = []byte("data")
|
||||
retryKey = []byte("retry")
|
||||
)
|
||||
|
||||
return func(yield func(Event, error) bool) {
|
||||
// iterate event from the wire.
|
||||
// https://developer.mozilla.org/en-US/docs/Web/API/Server-sent_events/Using_server-sent_events#examples
|
||||
//
|
||||
// - `key: value` line records.
|
||||
// - Consecutive `data: ...` fields are joined with newlines.
|
||||
// - Unrecognized fields are ignored. Since we only care about 'event', 'id', and
|
||||
// 'data', these are the only three we consider.
|
||||
// - Lines starting with ":" are ignored.
|
||||
// - Records are terminated with two consecutive newlines.
|
||||
var (
|
||||
evt Event
|
||||
dataBuf *bytes.Buffer // if non-nil, preceding field was also data
|
||||
)
|
||||
yieldEvent := func() bool {
|
||||
if dataBuf != nil {
|
||||
evt.Data = dataBuf.Bytes()
|
||||
dataBuf = nil
|
||||
}
|
||||
if evt.Empty() {
|
||||
return true
|
||||
}
|
||||
if !yield(evt, nil) {
|
||||
return false
|
||||
}
|
||||
evt = Event{}
|
||||
return true
|
||||
}
|
||||
for {
|
||||
line, err := reader.ReadBytes('\n')
|
||||
if err != nil && !errors.Is(err, io.EOF) {
|
||||
yield(Event{}, fmt.Errorf("error reading event: %v", err))
|
||||
return
|
||||
}
|
||||
line = bytes.TrimRight(line, "\r\n")
|
||||
isEOF := errors.Is(err, io.EOF)
|
||||
|
||||
if len(line) == 0 {
|
||||
if !yieldEvent() {
|
||||
return
|
||||
}
|
||||
if isEOF {
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
before, after, found := bytes.Cut(line, []byte{':'})
|
||||
if !found {
|
||||
yield(Event{}, fmt.Errorf("%w: malformed line in SSE stream: %q", errMalformedEvent, string(line)))
|
||||
return
|
||||
}
|
||||
switch {
|
||||
case bytes.Equal(before, eventKey):
|
||||
evt.Name = strings.TrimSpace(string(after))
|
||||
case bytes.Equal(before, idKey):
|
||||
evt.ID = strings.TrimSpace(string(after))
|
||||
case bytes.Equal(before, retryKey):
|
||||
evt.Retry = strings.TrimSpace(string(after))
|
||||
case bytes.Equal(before, dataKey):
|
||||
data := bytes.TrimSpace(after)
|
||||
if dataBuf == nil {
|
||||
dataBuf = new(bytes.Buffer)
|
||||
} else {
|
||||
dataBuf.WriteByte('\n')
|
||||
}
|
||||
dataBuf.Write(data)
|
||||
}
|
||||
|
||||
if isEOF {
|
||||
yieldEvent()
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// An EventStore tracks data for SSE streams.
|
||||
// A single EventStore suffices for all sessions, since session IDs are
|
||||
// globally unique. So one EventStore can be created per process, for
|
||||
// all Servers in the process.
|
||||
// Such a store is able to bound resource usage for the entire process.
|
||||
//
|
||||
// All of an EventStore's methods must be safe for use by multiple goroutines.
|
||||
type EventStore interface {
|
||||
// Open is called when a new stream is created. It may be used to ensure that
|
||||
// the underlying data structure for the stream is initialized, making it
|
||||
// ready to store and replay event streams.
|
||||
Open(_ context.Context, sessionID, streamID string) error
|
||||
|
||||
// Append appends data for an outgoing event to given stream, which is part of the
|
||||
// given session.
|
||||
Append(_ context.Context, sessionID, streamID string, data []byte) error
|
||||
|
||||
// After returns an iterator over the data for the given session and stream, beginning
|
||||
// just after the given index.
|
||||
//
|
||||
// Once the iterator yields a non-nil error, it will stop.
|
||||
// After's iterator must return an error immediately if any data after index was
|
||||
// dropped; it must not return partial results.
|
||||
// The stream must have been opened previously (see [EventStore.Open]).
|
||||
After(_ context.Context, sessionID, streamID string, index int) iter.Seq2[[]byte, error]
|
||||
|
||||
// SessionClosed informs the store that the given session is finished, along
|
||||
// with all of its streams.
|
||||
//
|
||||
// A store cannot rely on this method being called for cleanup. It should institute
|
||||
// additional mechanisms, such as timeouts, to reclaim storage.
|
||||
SessionClosed(_ context.Context, sessionID string) error
|
||||
|
||||
// There is no StreamClosed method. A server doesn't know when a stream is finished, because
|
||||
// the client can always send a GET with a Last-Event-ID referring to the stream.
|
||||
}
|
||||
|
||||
// A dataList is a list of []byte.
|
||||
// The zero dataList is ready to use.
|
||||
type dataList struct {
|
||||
size int // total size of data bytes
|
||||
first int // the stream index of the first element in data
|
||||
data [][]byte
|
||||
}
|
||||
|
||||
func (dl *dataList) appendData(d []byte) {
|
||||
// Empty data consumes memory but doesn't increment size. However, it should
|
||||
// be rare.
|
||||
dl.data = append(dl.data, d)
|
||||
dl.size += len(d)
|
||||
}
|
||||
|
||||
// removeFirst removes the first data item in dl, returning the size of the item.
|
||||
// It panics if dl is empty.
|
||||
func (dl *dataList) removeFirst() int {
|
||||
if len(dl.data) == 0 {
|
||||
panic("empty dataList")
|
||||
}
|
||||
r := len(dl.data[0])
|
||||
dl.size -= r
|
||||
dl.data[0] = nil // help GC
|
||||
dl.data = dl.data[1:]
|
||||
dl.first++
|
||||
return r
|
||||
}
|
||||
|
||||
// A MemoryEventStore is an [EventStore] backed by memory.
|
||||
type MemoryEventStore struct {
|
||||
mu sync.Mutex
|
||||
maxBytes int // max total size of all data
|
||||
nBytes int // current total size of all data
|
||||
store map[string]map[string]*dataList // session ID -> stream ID -> *dataList
|
||||
}
|
||||
|
||||
// MemoryEventStoreOptions are options for a [MemoryEventStore].
|
||||
type MemoryEventStoreOptions struct{}
|
||||
|
||||
// MaxBytes returns the maximum number of bytes that the store will retain before
|
||||
// purging data.
|
||||
func (s *MemoryEventStore) MaxBytes() int {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.maxBytes
|
||||
}
|
||||
|
||||
// SetMaxBytes sets the maximum number of bytes the store will retain before purging
|
||||
// data. The argument must not be negative. If it is zero, a suitable default will be used.
|
||||
// SetMaxBytes can be called at any time. The size of the store will be adjusted
|
||||
// immediately.
|
||||
func (s *MemoryEventStore) SetMaxBytes(n int) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
switch {
|
||||
case n < 0:
|
||||
panic("negative argument")
|
||||
case n == 0:
|
||||
s.maxBytes = defaultMaxBytes
|
||||
default:
|
||||
s.maxBytes = n
|
||||
}
|
||||
s.purge()
|
||||
}
|
||||
|
||||
const defaultMaxBytes = 10 << 20 // 10 MiB
|
||||
|
||||
// NewMemoryEventStore creates a [MemoryEventStore] with the default value
|
||||
// for MaxBytes.
|
||||
func NewMemoryEventStore(opts *MemoryEventStoreOptions) *MemoryEventStore {
|
||||
return &MemoryEventStore{
|
||||
maxBytes: defaultMaxBytes,
|
||||
store: make(map[string]map[string]*dataList),
|
||||
}
|
||||
}
|
||||
|
||||
// Open implements [EventStore.Open]. It ensures that the underlying data
|
||||
// structures for the given session are initialized and ready for use.
|
||||
func (s *MemoryEventStore) Open(_ context.Context, sessionID, streamID string) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.init(sessionID, streamID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// init is an internal helper function that ensures the nested map structure for a
|
||||
// given sessionID and streamID exists, creating it if necessary. It returns the
|
||||
// dataList associated with the specified IDs.
|
||||
// Requires s.mu.
|
||||
func (s *MemoryEventStore) init(sessionID, streamID string) *dataList {
|
||||
streamMap, ok := s.store[sessionID]
|
||||
if !ok {
|
||||
streamMap = make(map[string]*dataList)
|
||||
s.store[sessionID] = streamMap
|
||||
}
|
||||
dl, ok := streamMap[streamID]
|
||||
if !ok {
|
||||
dl = &dataList{}
|
||||
streamMap[streamID] = dl
|
||||
}
|
||||
return dl
|
||||
}
|
||||
|
||||
// Append implements [EventStore.Append] by recording data in memory.
|
||||
func (s *MemoryEventStore) Append(_ context.Context, sessionID, streamID string, data []byte) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
dl := s.init(sessionID, streamID)
|
||||
// Purge before adding, so at least the current data item will be present.
|
||||
// (That could result in nBytes > maxBytes, but we'll live with that.)
|
||||
s.purge()
|
||||
dl.appendData(data)
|
||||
s.nBytes += len(data)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ErrEventsPurged is the error that [EventStore.After] should return if the event just after the
|
||||
// index is no longer available.
|
||||
var ErrEventsPurged = errors.New("data purged")
|
||||
|
||||
// errMalformedEvent is returned when an SSE event cannot be parsed due to format violations.
|
||||
// This is a hard error indicating corrupted data or protocol violations, as opposed to
|
||||
// transient I/O errors which may be retryable.
|
||||
var errMalformedEvent = errors.New("malformed event")
|
||||
|
||||
// After implements [EventStore.After].
|
||||
func (s *MemoryEventStore) After(_ context.Context, sessionID, streamID string, index int) iter.Seq2[[]byte, error] {
|
||||
// Return the data items to yield.
|
||||
// We must copy, because dataList.removeFirst nils out slice elements.
|
||||
copyData := func() ([][]byte, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
streamMap, ok := s.store[sessionID]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("MemoryEventStore.After: unknown session ID %q", sessionID)
|
||||
}
|
||||
dl, ok := streamMap[streamID]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("MemoryEventStore.After: unknown stream ID %v in session %q", streamID, sessionID)
|
||||
}
|
||||
start := index + 1
|
||||
if dl.first > start {
|
||||
return nil, fmt.Errorf("MemoryEventStore.After: index %d, stream ID %v, session %q: %w",
|
||||
index, streamID, sessionID, ErrEventsPurged)
|
||||
}
|
||||
return slices.Clone(dl.data[start-dl.first:]), nil
|
||||
}
|
||||
|
||||
return func(yield func([]byte, error) bool) {
|
||||
ds, err := copyData()
|
||||
if err != nil {
|
||||
yield(nil, err)
|
||||
return
|
||||
}
|
||||
for _, d := range ds {
|
||||
if !yield(d, nil) {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// SessionClosed implements [EventStore.SessionClosed].
|
||||
func (s *MemoryEventStore) SessionClosed(_ context.Context, sessionID string) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
for _, dl := range s.store[sessionID] {
|
||||
s.nBytes -= dl.size
|
||||
}
|
||||
delete(s.store, sessionID)
|
||||
s.validate()
|
||||
return nil
|
||||
}
|
||||
|
||||
// purge removes data until no more than s.maxBytes bytes are in use.
|
||||
// It must be called with s.mu held.
|
||||
func (s *MemoryEventStore) purge() {
|
||||
// Remove the first element of every dataList until below the max.
|
||||
for s.nBytes > s.maxBytes {
|
||||
changed := false
|
||||
for _, sm := range s.store {
|
||||
for _, dl := range sm {
|
||||
if dl.size > 0 {
|
||||
r := dl.removeFirst()
|
||||
if r > 0 {
|
||||
changed = true
|
||||
s.nBytes -= r
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if !changed {
|
||||
panic("no progress during purge")
|
||||
}
|
||||
}
|
||||
s.validate()
|
||||
}
|
||||
|
||||
// validate checks that the store's data structures are valid.
|
||||
// It must be called with s.mu held.
|
||||
func (s *MemoryEventStore) validate() {
|
||||
if !validateMemoryEventStore {
|
||||
return
|
||||
}
|
||||
// Check that we're accounting for the size correctly.
|
||||
n := 0
|
||||
for _, sm := range s.store {
|
||||
for _, dl := range sm {
|
||||
for _, d := range dl.data {
|
||||
n += len(d)
|
||||
}
|
||||
}
|
||||
}
|
||||
if n != s.nBytes {
|
||||
panic("sizes don't add up")
|
||||
}
|
||||
}
|
||||
|
||||
// debugString returns a string containing the state of s.
|
||||
// Used in tests.
|
||||
func (s *MemoryEventStore) debugString() string {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
var b strings.Builder
|
||||
for i, sess := range slices.Sorted(maps.Keys(s.store)) {
|
||||
if i > 0 {
|
||||
fmt.Fprintf(&b, "; ")
|
||||
}
|
||||
sm := s.store[sess]
|
||||
for i, sid := range slices.Sorted(maps.Keys(sm)) {
|
||||
if i > 0 {
|
||||
fmt.Fprintf(&b, "; ")
|
||||
}
|
||||
dl := sm[sid]
|
||||
fmt.Fprintf(&b, "%s %s first=%d", sess, sid, dl.first)
|
||||
for _, d := range dl.data {
|
||||
fmt.Fprintf(&b, " %s", d)
|
||||
}
|
||||
}
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
+114
@@ -0,0 +1,114 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"iter"
|
||||
"maps"
|
||||
"slices"
|
||||
)
|
||||
|
||||
// This file contains implementations that are common to all features.
|
||||
// A feature is an item provided to a peer. In the 2025-03-26 spec,
|
||||
// the features are prompt, tool, resource and root.
|
||||
|
||||
// A featureSet is a collection of features of type T.
|
||||
// Every feature has a unique ID, and the spec never mentions
|
||||
// an ordering for the List calls, so what it calls a "list" is actually a set.
|
||||
//
|
||||
// An alternative implementation would use an ordered map, but that's probably
|
||||
// not necessary as adds and removes are rare, and usually batched.
|
||||
type featureSet[T any] struct {
|
||||
uniqueID func(T) string
|
||||
features map[string]T
|
||||
sortedKeys []string // lazily computed; nil after add or remove
|
||||
}
|
||||
|
||||
// newFeatureSet creates a new featureSet for features of type T.
|
||||
// The argument function should return the unique ID for a single feature.
|
||||
func newFeatureSet[T any](uniqueIDFunc func(T) string) *featureSet[T] {
|
||||
return &featureSet[T]{
|
||||
uniqueID: uniqueIDFunc,
|
||||
features: make(map[string]T),
|
||||
}
|
||||
}
|
||||
|
||||
// add adds each feature to the set if it is not present,
|
||||
// or replaces an existing feature.
|
||||
func (s *featureSet[T]) add(fs ...T) {
|
||||
for _, f := range fs {
|
||||
s.features[s.uniqueID(f)] = f
|
||||
}
|
||||
s.sortedKeys = nil
|
||||
}
|
||||
|
||||
// remove removes all features with the given uids from the set if present,
|
||||
// and returns whether any were removed.
|
||||
// It is not an error to remove a nonexistent feature.
|
||||
func (s *featureSet[T]) remove(uids ...string) bool {
|
||||
changed := false
|
||||
for _, uid := range uids {
|
||||
if _, ok := s.features[uid]; ok {
|
||||
changed = true
|
||||
delete(s.features, uid)
|
||||
}
|
||||
}
|
||||
if changed {
|
||||
s.sortedKeys = nil
|
||||
}
|
||||
return changed
|
||||
}
|
||||
|
||||
// get returns the feature with the given uid.
|
||||
// If there is none, it returns zero, false.
|
||||
func (s *featureSet[T]) get(uid string) (T, bool) {
|
||||
t, ok := s.features[uid]
|
||||
return t, ok
|
||||
}
|
||||
|
||||
// len returns the number of features in the set.
|
||||
func (s *featureSet[T]) len() int { return len(s.features) }
|
||||
|
||||
// all returns an iterator over of all the features in the set
|
||||
// sorted by unique ID.
|
||||
func (s *featureSet[T]) all() iter.Seq[T] {
|
||||
s.sortKeys()
|
||||
return func(yield func(T) bool) {
|
||||
s.yieldFrom(0, yield)
|
||||
}
|
||||
}
|
||||
|
||||
// above returns an iterator over features in the set whose unique IDs are
|
||||
// greater than `uid`, in ascending ID order.
|
||||
func (s *featureSet[T]) above(uid string) iter.Seq[T] {
|
||||
s.sortKeys()
|
||||
index, found := slices.BinarySearch(s.sortedKeys, uid)
|
||||
if found {
|
||||
index++
|
||||
}
|
||||
return func(yield func(T) bool) {
|
||||
s.yieldFrom(index, yield)
|
||||
}
|
||||
}
|
||||
|
||||
// sortKeys is a helper that maintains a sorted list of feature IDs. It
|
||||
// computes this list lazily upon its first call after a modification, or
|
||||
// if it's nil.
|
||||
func (s *featureSet[T]) sortKeys() {
|
||||
if s.sortedKeys != nil {
|
||||
return
|
||||
}
|
||||
s.sortedKeys = slices.Sorted(maps.Keys(s.features))
|
||||
}
|
||||
|
||||
// yieldFrom is a helper that iterates over the features in the set,
|
||||
// starting at the given index, and calls the yield function for each one.
|
||||
func (s *featureSet[T]) yieldFrom(index int, yield func(T) bool) {
|
||||
for i := index; i < len(s.sortedKeys); i++ {
|
||||
if !yield(s.features[s.sortedKeys[i]]) {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
+201
@@ -0,0 +1,201 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"cmp"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"slices"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Logging levels.
|
||||
const (
|
||||
LevelDebug = slog.LevelDebug
|
||||
LevelInfo = slog.LevelInfo
|
||||
LevelNotice = (slog.LevelInfo + slog.LevelWarn) / 2
|
||||
LevelWarning = slog.LevelWarn
|
||||
LevelError = slog.LevelError
|
||||
LevelCritical = slog.LevelError + 4
|
||||
LevelAlert = slog.LevelError + 8
|
||||
LevelEmergency = slog.LevelError + 12
|
||||
)
|
||||
|
||||
var slogToMCP = map[slog.Level]LoggingLevel{
|
||||
LevelDebug: "debug",
|
||||
LevelInfo: "info",
|
||||
LevelNotice: "notice",
|
||||
LevelWarning: "warning",
|
||||
LevelError: "error",
|
||||
LevelCritical: "critical",
|
||||
LevelAlert: "alert",
|
||||
LevelEmergency: "emergency",
|
||||
}
|
||||
|
||||
var mcpToSlog = make(map[LoggingLevel]slog.Level)
|
||||
|
||||
func init() {
|
||||
for sl, ml := range slogToMCP {
|
||||
mcpToSlog[ml] = sl
|
||||
}
|
||||
}
|
||||
|
||||
func slogLevelToMCP(sl slog.Level) LoggingLevel {
|
||||
if ml, ok := slogToMCP[sl]; ok {
|
||||
return ml
|
||||
}
|
||||
return "debug" // for lack of a better idea
|
||||
}
|
||||
|
||||
func mcpLevelToSlog(ll LoggingLevel) slog.Level {
|
||||
if sl, ok := mcpToSlog[ll]; ok {
|
||||
return sl
|
||||
}
|
||||
// TODO: is there a better default?
|
||||
return LevelDebug
|
||||
}
|
||||
|
||||
// compareLevels behaves like [cmp.Compare] for [LoggingLevel]s.
|
||||
func compareLevels(l1, l2 LoggingLevel) int {
|
||||
return cmp.Compare(mcpLevelToSlog(l1), mcpLevelToSlog(l2))
|
||||
}
|
||||
|
||||
// LoggingHandlerOptions are options for a LoggingHandler.
|
||||
type LoggingHandlerOptions struct {
|
||||
// The value for the "logger" field of logging notifications.
|
||||
LoggerName string
|
||||
// Limits the rate at which log messages are sent.
|
||||
// Excess messages are dropped.
|
||||
// If zero, there is no rate limiting.
|
||||
MinInterval time.Duration
|
||||
}
|
||||
|
||||
// A LoggingHandler is a [slog.Handler] for MCP.
|
||||
type LoggingHandler struct {
|
||||
opts LoggingHandlerOptions
|
||||
ss *ServerSession
|
||||
// Ensures that the buffer reset is atomic with the write (see Handle).
|
||||
// A pointer so that clones share the mutex. See
|
||||
// https://github.com/golang/example/blob/master/slog-handler-guide/README.md#getting-the-mutex-right.
|
||||
mu *sync.Mutex
|
||||
lastMessageSent time.Time // for rate-limiting
|
||||
buf *bytes.Buffer
|
||||
handler slog.Handler
|
||||
}
|
||||
|
||||
// ensureLogger returns l if non-nil, otherwise a discard logger.
|
||||
func ensureLogger(l *slog.Logger) *slog.Logger {
|
||||
if l != nil {
|
||||
return l
|
||||
}
|
||||
return slog.New(slog.DiscardHandler)
|
||||
}
|
||||
|
||||
// NewLoggingHandler creates a [LoggingHandler] that logs to the given [ServerSession] using a
|
||||
// [slog.JSONHandler].
|
||||
func NewLoggingHandler(ss *ServerSession, opts *LoggingHandlerOptions) *LoggingHandler {
|
||||
var buf bytes.Buffer
|
||||
jsonHandler := slog.NewJSONHandler(&buf, &slog.HandlerOptions{
|
||||
ReplaceAttr: func(_ []string, a slog.Attr) slog.Attr {
|
||||
// Remove level: it appears in LoggingMessageParams.
|
||||
if a.Key == slog.LevelKey {
|
||||
return slog.Attr{}
|
||||
}
|
||||
return a
|
||||
},
|
||||
})
|
||||
lh := &LoggingHandler{
|
||||
ss: ss,
|
||||
mu: new(sync.Mutex),
|
||||
buf: &buf,
|
||||
handler: jsonHandler,
|
||||
}
|
||||
if opts != nil {
|
||||
lh.opts = *opts
|
||||
}
|
||||
return lh
|
||||
}
|
||||
|
||||
// Enabled implements [slog.Handler.Enabled] by comparing level to the [ServerSession]'s level.
|
||||
func (h *LoggingHandler) Enabled(ctx context.Context, level slog.Level) bool {
|
||||
// This is also checked in ServerSession.LoggingMessage, so checking it here
|
||||
// is just an optimization that skips building the JSON.
|
||||
h.ss.mu.Lock()
|
||||
mcpLevel := h.ss.state.LogLevel
|
||||
h.ss.mu.Unlock()
|
||||
return level >= mcpLevelToSlog(mcpLevel)
|
||||
}
|
||||
|
||||
// WithAttrs implements [slog.Handler.WithAttrs].
|
||||
func (h *LoggingHandler) WithAttrs(as []slog.Attr) slog.Handler {
|
||||
h2 := *h
|
||||
h2.handler = h.handler.WithAttrs(as)
|
||||
return &h2
|
||||
}
|
||||
|
||||
// WithGroup implements [slog.Handler.WithGroup].
|
||||
func (h *LoggingHandler) WithGroup(name string) slog.Handler {
|
||||
h2 := *h
|
||||
h2.handler = h.handler.WithGroup(name)
|
||||
return &h2
|
||||
}
|
||||
|
||||
// Handle implements [slog.Handler.Handle] by writing the Record to a JSONHandler,
|
||||
// then calling [ServerSession.LoggingMessage] with the result.
|
||||
func (h *LoggingHandler) Handle(ctx context.Context, r slog.Record) error {
|
||||
err := h.handle(ctx, r)
|
||||
// TODO(jba): find a way to surface the error.
|
||||
// The return value will probably be ignored.
|
||||
return err
|
||||
}
|
||||
|
||||
func (h *LoggingHandler) handle(ctx context.Context, r slog.Record) error {
|
||||
// Observe the rate limit.
|
||||
// TODO(jba): use golang.org/x/time/rate.
|
||||
h.mu.Lock()
|
||||
skip := time.Since(h.lastMessageSent) < h.opts.MinInterval
|
||||
h.mu.Unlock()
|
||||
if skip {
|
||||
return nil
|
||||
}
|
||||
|
||||
var err error
|
||||
var data json.RawMessage
|
||||
// Make the buffer reset atomic with the record write.
|
||||
// We are careful here in the unlikely event that the handler panics.
|
||||
// We don't want to hold the lock for the entire function, because Notify is
|
||||
// an I/O operation.
|
||||
// This can result in out-of-order delivery.
|
||||
func() {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
h.buf.Reset()
|
||||
err = h.handler.Handle(ctx, r)
|
||||
// Clone the buffer as Bytes() references the internal buffer.
|
||||
data = json.RawMessage(slices.Clone(h.buf.Bytes()))
|
||||
}()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
h.mu.Lock()
|
||||
h.lastMessageSent = time.Now()
|
||||
h.mu.Unlock()
|
||||
|
||||
params := &LoggingMessageParams{
|
||||
Logger: h.opts.LoggerName,
|
||||
Level: slogLevelToMCP(r.Level),
|
||||
Data: data,
|
||||
}
|
||||
// We pass the argument context to Notify, even though slog.Handler.Handle's
|
||||
// documentation says not to.
|
||||
// In this case logging is a service to clients, not a means for debugging the
|
||||
// server, so we want to cancel the log message.
|
||||
return h.ss.Log(ctx, params)
|
||||
}
|
||||
+88
@@ -0,0 +1,88 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// The mcp package provides an SDK for writing model context protocol clients
|
||||
// and servers.
|
||||
//
|
||||
// To get started, create either a [Client] or [Server], add features to it
|
||||
// using `AddXXX` functions, and connect it to a peer using a [Transport].
|
||||
//
|
||||
// For example, to run a simple server on the [StdioTransport]:
|
||||
//
|
||||
// server := mcp.NewServer(&mcp.Implementation{Name: "greeter"}, nil)
|
||||
//
|
||||
// // Using the generic AddTool automatically populates the the input and output
|
||||
// // schema of the tool.
|
||||
// type args struct {
|
||||
// Name string `json:"name" jsonschema:"the person to greet"`
|
||||
// }
|
||||
// mcp.AddTool(server, &mcp.Tool{
|
||||
// Name: "greet",
|
||||
// Description: "say hi",
|
||||
// }, func(ctx context.Context, req *mcp.CallToolRequest, args args) (*mcp.CallToolResult, any, error) {
|
||||
// return &mcp.CallToolResult{
|
||||
// Content: []mcp.Content{
|
||||
// &mcp.TextContent{Text: "Hi " + args.Name},
|
||||
// },
|
||||
// }, nil, nil
|
||||
// })
|
||||
//
|
||||
// // Run the server on the stdio transport.
|
||||
// if err := server.Run(context.Background(), &mcp.StdioTransport{}); err != nil {
|
||||
// log.Printf("Server failed: %v", err)
|
||||
// }
|
||||
//
|
||||
// To connect to this server, use the [CommandTransport]:
|
||||
//
|
||||
// client := mcp.NewClient(&mcp.Implementation{Name: "mcp-client", Version: "v1.0.0"}, nil)
|
||||
// transport := &mcp.CommandTransport{Command: exec.Command("myserver")}
|
||||
// session, err := client.Connect(ctx, transport, nil)
|
||||
// if err != nil {
|
||||
// log.Fatal(err)
|
||||
// }
|
||||
// defer session.Close()
|
||||
//
|
||||
// params := &mcp.CallToolParams{
|
||||
// Name: "greet",
|
||||
// Arguments: map[string]any{"name": "you"},
|
||||
// }
|
||||
// res, err := session.CallTool(ctx, params)
|
||||
// if err != nil {
|
||||
// log.Fatalf("CallTool failed: %v", err)
|
||||
// }
|
||||
//
|
||||
// # Clients, servers, and sessions
|
||||
//
|
||||
// In this SDK, both a [Client] and [Server] may handle many concurrent
|
||||
// connections. Each time a client or server is connected to a peer using a
|
||||
// [Transport], it creates a new session (either a [ClientSession] or
|
||||
// [ServerSession]):
|
||||
//
|
||||
// Client Server
|
||||
// ⇅ (jsonrpc2) ⇅
|
||||
// ClientSession ⇄ Client Transport ⇄ Server Transport ⇄ ServerSession
|
||||
//
|
||||
// The session types expose an API to interact with its peer. For example,
|
||||
// [ClientSession.CallTool] or [ServerSession.ListRoots].
|
||||
//
|
||||
// # Adding features
|
||||
//
|
||||
// Add MCP servers to your Client or Server using AddXXX methods (for example
|
||||
// [Client.AddRoot] or [Server.AddPrompt]). If any peers are connected when
|
||||
// AddXXX is called, they will receive a corresponding change notification
|
||||
// (for example notifications/roots/list_changed).
|
||||
//
|
||||
// Adding tools is special: tools may be bound to ordinary Go functions by
|
||||
// using the top-level generic [AddTool] function, which allows specifying an
|
||||
// input and output type. When AddTool is used, the tool's input schema and
|
||||
// output schema are automatically populated, and inputs are automatically
|
||||
// validated. As a special case, if the output type is 'any', no output schema
|
||||
// is generated.
|
||||
//
|
||||
// func double(_ context.Context, _ *mcp.CallToolRequest, in In) (*mcp.CallToolResult, Out, error) {
|
||||
// return nil, Out{Answer: 2*in.Number}, nil
|
||||
// }
|
||||
// ...
|
||||
// mcp.AddTool(server, &mcp.Tool{Name: "double"}, double)
|
||||
package mcp
|
||||
+17
@@ -0,0 +1,17 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
)
|
||||
|
||||
// A PromptHandler handles a call to prompts/get.
|
||||
type PromptHandler func(context.Context, *GetPromptRequest) (*GetPromptResult, error)
|
||||
|
||||
type serverPrompt struct {
|
||||
prompt *Prompt
|
||||
handler PromptHandler
|
||||
}
|
||||
+1622
File diff suppressed because it is too large
Load Diff
+39
@@ -0,0 +1,39 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// This file holds the request types.
|
||||
|
||||
package mcp
|
||||
|
||||
type (
|
||||
CallToolRequest = ServerRequest[*CallToolParamsRaw]
|
||||
CompleteRequest = ServerRequest[*CompleteParams]
|
||||
GetPromptRequest = ServerRequest[*GetPromptParams]
|
||||
InitializedRequest = ServerRequest[*InitializedParams]
|
||||
ListPromptsRequest = ServerRequest[*ListPromptsParams]
|
||||
ListResourcesRequest = ServerRequest[*ListResourcesParams]
|
||||
ListResourceTemplatesRequest = ServerRequest[*ListResourceTemplatesParams]
|
||||
ListToolsRequest = ServerRequest[*ListToolsParams]
|
||||
ProgressNotificationServerRequest = ServerRequest[*ProgressNotificationParams]
|
||||
ReadResourceRequest = ServerRequest[*ReadResourceParams]
|
||||
RootsListChangedRequest = ServerRequest[*RootsListChangedParams]
|
||||
SubscribeRequest = ServerRequest[*SubscribeParams]
|
||||
UnsubscribeRequest = ServerRequest[*UnsubscribeParams]
|
||||
)
|
||||
|
||||
type (
|
||||
CreateMessageRequest = ClientRequest[*CreateMessageParams]
|
||||
CreateMessageWithToolsRequest = ClientRequest[*CreateMessageWithToolsParams]
|
||||
ElicitRequest = ClientRequest[*ElicitParams]
|
||||
initializedClientRequest = ClientRequest[*InitializedParams]
|
||||
InitializeRequest = ClientRequest[*InitializeParams]
|
||||
ListRootsRequest = ClientRequest[*ListRootsParams]
|
||||
LoggingMessageRequest = ClientRequest[*LoggingMessageParams]
|
||||
ProgressNotificationClientRequest = ClientRequest[*ProgressNotificationParams]
|
||||
PromptListChangedRequest = ClientRequest[*PromptListChangedParams]
|
||||
ResourceListChangedRequest = ClientRequest[*ResourceListChangedParams]
|
||||
ResourceUpdatedNotificationRequest = ClientRequest[*ResourceUpdatedNotificationParams]
|
||||
ToolListChangedRequest = ClientRequest[*ToolListChangedParams]
|
||||
ElicitationCompleteNotificationRequest = ClientRequest[*ElicitationCompleteParams]
|
||||
)
|
||||
+181
@@ -0,0 +1,181 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/internal/util"
|
||||
"github.com/modelcontextprotocol/go-sdk/jsonrpc"
|
||||
"github.com/yosida95/uritemplate/v3"
|
||||
)
|
||||
|
||||
// A serverResource associates a Resource with its handler.
|
||||
type serverResource struct {
|
||||
resource *Resource
|
||||
handler ResourceHandler
|
||||
}
|
||||
|
||||
// A serverResourceTemplate associates a ResourceTemplate with its handler.
|
||||
type serverResourceTemplate struct {
|
||||
resourceTemplate *ResourceTemplate
|
||||
handler ResourceHandler
|
||||
}
|
||||
|
||||
// A ResourceHandler is a function that reads a resource.
|
||||
// It will be called when the client calls [ClientSession.ReadResource].
|
||||
// If it cannot find the resource, it should return the result of calling [ResourceNotFoundError].
|
||||
type ResourceHandler func(context.Context, *ReadResourceRequest) (*ReadResourceResult, error)
|
||||
|
||||
// ResourceNotFoundError returns an error indicating that a resource being read could
|
||||
// not be found.
|
||||
func ResourceNotFoundError(uri string) error {
|
||||
return &jsonrpc.Error{
|
||||
Code: CodeResourceNotFound,
|
||||
Message: "Resource not found",
|
||||
Data: json.RawMessage(fmt.Sprintf(`{"uri":%q}`, uri)),
|
||||
}
|
||||
}
|
||||
|
||||
// readFileResource reads from the filesystem at a URI relative to dirFilepath, respecting
|
||||
// the roots.
|
||||
// dirFilepath and rootFilepaths are absolute filesystem paths.
|
||||
func readFileResource(rawURI, dirFilepath string, rootFilepaths []string) ([]byte, error) {
|
||||
uriFilepath, err := computeURIFilepath(rawURI, dirFilepath, rootFilepaths)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var data []byte
|
||||
err = withFile(dirFilepath, uriFilepath, func(f *os.File) error {
|
||||
var err error
|
||||
data, err = io.ReadAll(f)
|
||||
return err
|
||||
})
|
||||
if os.IsNotExist(err) {
|
||||
err = ResourceNotFoundError(rawURI)
|
||||
}
|
||||
return data, err
|
||||
}
|
||||
|
||||
// computeURIFilepath returns a path relative to dirFilepath.
|
||||
// The dirFilepath and rootFilepaths are absolute file paths.
|
||||
func computeURIFilepath(rawURI, dirFilepath string, rootFilepaths []string) (string, error) {
|
||||
// We use "file path" to mean a filesystem path.
|
||||
uri, err := url.Parse(rawURI)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if uri.Scheme != "file" {
|
||||
return "", fmt.Errorf("URI is not a file: %s", uri)
|
||||
}
|
||||
if uri.Path == "" {
|
||||
// A more specific error than the one below, to catch the
|
||||
// common mistake "file://foo".
|
||||
return "", errors.New("empty path")
|
||||
}
|
||||
// The URI's path is interpreted relative to dirFilepath, and in the local filesystem.
|
||||
// It must not try to escape its directory.
|
||||
uriFilepathRel, err := filepath.Localize(strings.TrimPrefix(uri.Path, "/"))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%q cannot be localized: %w", uriFilepathRel, err)
|
||||
}
|
||||
|
||||
// Check roots, if there are any.
|
||||
if len(rootFilepaths) > 0 {
|
||||
// To check against the roots, we need an absolute file path, not relative to the directory.
|
||||
// uriFilepath is local, so the joined path is under dirFilepath.
|
||||
uriFilepathAbs := filepath.Join(dirFilepath, uriFilepathRel)
|
||||
rootOK := false
|
||||
// Check that the requested file path is under some root.
|
||||
// Since both paths are absolute, that's equivalent to filepath.Rel constructing
|
||||
// a local path.
|
||||
for _, rootFilepathAbs := range rootFilepaths {
|
||||
if rel, err := filepath.Rel(rootFilepathAbs, uriFilepathAbs); err == nil && filepath.IsLocal(rel) {
|
||||
rootOK = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !rootOK {
|
||||
return "", fmt.Errorf("URI path %q is not under any root", uriFilepathAbs)
|
||||
}
|
||||
}
|
||||
return uriFilepathRel, nil
|
||||
}
|
||||
|
||||
// withFile calls f on the file at join(dir, rel),
|
||||
// protecting against path traversal attacks.
|
||||
func withFile(dir, rel string, f func(*os.File) error) (err error) {
|
||||
r, err := os.OpenRoot(dir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer r.Close()
|
||||
file, err := r.Open(rel)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// Record error, in case f writes.
|
||||
defer func() { err = errors.Join(err, file.Close()) }()
|
||||
return f(file)
|
||||
}
|
||||
|
||||
// fileRoots transforms the Roots obtained from the client into absolute paths on
|
||||
// the local filesystem.
|
||||
// TODO(jba): expose this functionality to user ResourceHandlers,
|
||||
// so they don't have to repeat it.
|
||||
func fileRoots(rawRoots []*Root) ([]string, error) {
|
||||
var fileRoots []string
|
||||
for _, r := range rawRoots {
|
||||
fr, err := fileRoot(r)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
fileRoots = append(fileRoots, fr)
|
||||
}
|
||||
return fileRoots, nil
|
||||
}
|
||||
|
||||
// fileRoot returns the absolute path for Root.
|
||||
func fileRoot(root *Root) (_ string, err error) {
|
||||
defer util.Wrapf(&err, "root %q", root.URI)
|
||||
|
||||
// Convert to absolute file path.
|
||||
rurl, err := url.Parse(root.URI)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if rurl.Scheme != "file" {
|
||||
return "", errors.New("not a file URI")
|
||||
}
|
||||
if rurl.Path == "" {
|
||||
// A more specific error than the one below, to catch the
|
||||
// common mistake "file://foo".
|
||||
return "", errors.New("empty path")
|
||||
}
|
||||
// We don't want Localize here: we want an absolute path, which is not local.
|
||||
fileRoot := filepath.Clean(filepath.FromSlash(rurl.Path))
|
||||
if !filepath.IsAbs(fileRoot) {
|
||||
return "", errors.New("not an absolute path")
|
||||
}
|
||||
return fileRoot, nil
|
||||
}
|
||||
|
||||
// Matches reports whether the receiver's uri template matches the uri.
|
||||
func (sr *serverResourceTemplate) Matches(uri string) bool {
|
||||
tmpl, err := uritemplate.New(sr.resourceTemplate.URITemplate)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return tmpl.Regexp().MatchString(uri)
|
||||
}
|
||||
+69
@@ -0,0 +1,69 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"sync"
|
||||
|
||||
"github.com/google/jsonschema-go/jsonschema"
|
||||
)
|
||||
|
||||
// A SchemaCache caches JSON schemas to avoid repeated reflection and resolution.
|
||||
//
|
||||
// This is useful for stateless server deployments (one [Server] per request)
|
||||
// where tools are re-registered on every request. Without caching, each
|
||||
// [AddTool] call triggers expensive reflection-based schema generation.
|
||||
//
|
||||
// A SchemaCache is safe for concurrent use by multiple goroutines.
|
||||
//
|
||||
// # Trade-offs
|
||||
//
|
||||
// The cache is unbounded: it stores one entry per unique Go type or schema
|
||||
// pointer. For typical MCP servers with a fixed set of tools, memory usage
|
||||
// is negligible. However, if tool input types are generated dynamically,
|
||||
// the cache will grow without bound.
|
||||
//
|
||||
// The cache uses pointer identity for pre-defined schemas. If a schema's
|
||||
// contents change but the pointer remains the same, stale resolved schemas
|
||||
// may be returned. In practice, this is not an issue because tool schemas
|
||||
// are typically defined once at startup.
|
||||
type SchemaCache struct {
|
||||
byType sync.Map // reflect.Type -> *cachedSchema
|
||||
bySchema sync.Map // *jsonschema.Schema -> *jsonschema.Resolved
|
||||
}
|
||||
|
||||
type cachedSchema struct {
|
||||
schema *jsonschema.Schema
|
||||
resolved *jsonschema.Resolved
|
||||
}
|
||||
|
||||
// NewSchemaCache creates a new [SchemaCache].
|
||||
func NewSchemaCache() *SchemaCache {
|
||||
return &SchemaCache{}
|
||||
}
|
||||
|
||||
func (c *SchemaCache) getByType(t reflect.Type) (*jsonschema.Schema, *jsonschema.Resolved, bool) {
|
||||
if v, ok := c.byType.Load(t); ok {
|
||||
cs := v.(*cachedSchema)
|
||||
return cs.schema, cs.resolved, true
|
||||
}
|
||||
return nil, nil, false
|
||||
}
|
||||
|
||||
func (c *SchemaCache) setByType(t reflect.Type, schema *jsonschema.Schema, resolved *jsonschema.Resolved) {
|
||||
c.byType.Store(t, &cachedSchema{schema: schema, resolved: resolved})
|
||||
}
|
||||
|
||||
func (c *SchemaCache) getBySchema(schema *jsonschema.Schema) (*jsonschema.Resolved, bool) {
|
||||
if v, ok := c.bySchema.Load(schema); ok {
|
||||
return v.(*jsonschema.Resolved), true
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
func (c *SchemaCache) setBySchema(schema *jsonschema.Schema, resolved *jsonschema.Resolved) {
|
||||
c.bySchema.Store(schema, resolved)
|
||||
}
|
||||
+1595
File diff suppressed because it is too large
Load Diff
+29
@@ -0,0 +1,29 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package mcp
|
||||
|
||||
// hasSessionID is the interface which, if implemented by connections, informs
|
||||
// the session about their session ID.
|
||||
//
|
||||
// TODO(rfindley): remove SessionID methods from connections, when it doesn't
|
||||
// make sense. Or remove it from the Sessions entirely: why does it even need
|
||||
// to be exposed?
|
||||
type hasSessionID interface {
|
||||
SessionID() string
|
||||
}
|
||||
|
||||
// ServerSessionState is the state of a session.
|
||||
type ServerSessionState struct {
|
||||
// InitializeParams are the parameters from 'initialize'.
|
||||
InitializeParams *InitializeParams `json:"initializeParams"`
|
||||
|
||||
// InitializedParams are the parameters from 'notifications/initialized'.
|
||||
InitializedParams *InitializedParams `json:"initializedParams"`
|
||||
|
||||
// LogLevel is the logging level for the session.
|
||||
LogLevel LoggingLevel `json:"logLevel"`
|
||||
|
||||
// TODO: resource subscriptions
|
||||
}
|
||||
+611
@@ -0,0 +1,611 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// This file contains code shared between client and server, including
|
||||
// method handler and middleware definitions.
|
||||
//
|
||||
// Much of this is here so that we can factor out commonalities using
|
||||
// generics. If this becomes unwieldy, it can perhaps be simplified with
|
||||
// reflection.
|
||||
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"reflect"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/auth"
|
||||
internaljson "github.com/modelcontextprotocol/go-sdk/internal/json"
|
||||
"github.com/modelcontextprotocol/go-sdk/internal/jsonrpc2"
|
||||
"github.com/modelcontextprotocol/go-sdk/jsonrpc"
|
||||
)
|
||||
|
||||
const (
|
||||
// latestProtocolVersion is the latest protocol version that this version of
|
||||
// the SDK supports.
|
||||
//
|
||||
// It is the version that the client sends in the initialization request, and
|
||||
// the default version used by the server.
|
||||
latestProtocolVersion = protocolVersion20250618
|
||||
protocolVersion20251125 = "2025-11-25" // not yet released
|
||||
protocolVersion20250618 = "2025-06-18"
|
||||
protocolVersion20250326 = "2025-03-26"
|
||||
protocolVersion20241105 = "2024-11-05"
|
||||
)
|
||||
|
||||
var supportedProtocolVersions = []string{
|
||||
protocolVersion20251125,
|
||||
protocolVersion20250618,
|
||||
protocolVersion20250326,
|
||||
protocolVersion20241105,
|
||||
}
|
||||
|
||||
// negotiatedVersion returns the effective protocol version to use, given a
|
||||
// client version.
|
||||
func negotiatedVersion(clientVersion string) string {
|
||||
// In general, prefer to use the clientVersion, but if we don't support the
|
||||
// client's version, use the latest version.
|
||||
//
|
||||
// This handles the case where a new spec version is released, and the SDK
|
||||
// does not support it yet.
|
||||
if !slices.Contains(supportedProtocolVersions, clientVersion) {
|
||||
return latestProtocolVersion
|
||||
}
|
||||
return clientVersion
|
||||
}
|
||||
|
||||
// A MethodHandler handles MCP messages.
|
||||
// For methods, exactly one of the return values must be nil.
|
||||
// For notifications, both must be nil.
|
||||
type MethodHandler func(ctx context.Context, method string, req Request) (result Result, err error)
|
||||
|
||||
// A Session is either a [ClientSession] or a [ServerSession].
|
||||
type Session interface {
|
||||
// ID returns the session ID, or the empty string if there is none.
|
||||
ID() string
|
||||
|
||||
sendingMethodInfos() map[string]methodInfo
|
||||
receivingMethodInfos() map[string]methodInfo
|
||||
sendingMethodHandler() MethodHandler
|
||||
receivingMethodHandler() MethodHandler
|
||||
getConn() *jsonrpc2.Connection
|
||||
}
|
||||
|
||||
// Middleware is a function from [MethodHandler] to [MethodHandler].
|
||||
type Middleware func(MethodHandler) MethodHandler
|
||||
|
||||
// addMiddleware wraps the handler in the middleware functions.
|
||||
func addMiddleware(handlerp *MethodHandler, middleware []Middleware) {
|
||||
for _, m := range slices.Backward(middleware) {
|
||||
*handlerp = m(*handlerp)
|
||||
}
|
||||
}
|
||||
|
||||
func defaultSendingMethodHandler(ctx context.Context, method string, req Request) (Result, error) {
|
||||
info, ok := req.GetSession().sendingMethodInfos()[method]
|
||||
if !ok {
|
||||
// This can be called from user code, with an arbitrary value for method.
|
||||
return nil, jsonrpc2.ErrNotHandled
|
||||
}
|
||||
params := req.GetParams()
|
||||
if initParams, ok := params.(*InitializeParams); ok {
|
||||
// Fix the marshaling of initialize params, to work around #607.
|
||||
//
|
||||
// The initialize params we produce should never be nil, nor have nil
|
||||
// capabilities, so any panic here is a bug.
|
||||
params = initParams.toV2()
|
||||
}
|
||||
// Notifications don't have results.
|
||||
if strings.HasPrefix(method, "notifications/") {
|
||||
return nil, req.GetSession().getConn().Notify(ctx, method, params)
|
||||
}
|
||||
// Create the result to unmarshal into.
|
||||
// The concrete type of the result is the return type of the receiving function.
|
||||
res := info.newResult()
|
||||
if err := call(ctx, req.GetSession().getConn(), method, params, res); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return res, nil
|
||||
}
|
||||
|
||||
// Helper method to avoid typed nil.
|
||||
func orZero[T any, P *U, U any](p P) T {
|
||||
if p == nil {
|
||||
var zero T
|
||||
return zero
|
||||
}
|
||||
return any(p).(T)
|
||||
}
|
||||
|
||||
func handleNotify(ctx context.Context, method string, req Request) error {
|
||||
mh := req.GetSession().sendingMethodHandler()
|
||||
_, err := mh(ctx, method, req)
|
||||
return err
|
||||
}
|
||||
|
||||
func handleSend[R Result](ctx context.Context, method string, req Request) (R, error) {
|
||||
mh := req.GetSession().sendingMethodHandler()
|
||||
// mh might be user code, so ensure that it returns the right values for the jsonrpc2 protocol.
|
||||
res, err := mh(ctx, method, req)
|
||||
if err != nil {
|
||||
var z R
|
||||
return z, err
|
||||
}
|
||||
return res.(R), nil
|
||||
}
|
||||
|
||||
// defaultReceivingMethodHandler is the initial MethodHandler for servers and clients, before being wrapped by middleware.
|
||||
func defaultReceivingMethodHandler[S Session](ctx context.Context, method string, req Request) (Result, error) {
|
||||
info, ok := req.GetSession().receivingMethodInfos()[method]
|
||||
if !ok {
|
||||
// This can be called from user code, with an arbitrary value for method.
|
||||
return nil, jsonrpc2.ErrNotHandled
|
||||
}
|
||||
return info.handleMethod(ctx, method, req)
|
||||
}
|
||||
|
||||
func handleReceive[S Session](ctx context.Context, session S, jreq *jsonrpc.Request) (Result, error) {
|
||||
info, err := checkRequest(jreq, session.receivingMethodInfos())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
params, err := info.unmarshalParams(jreq.Params)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("handling '%s': %w", jreq.Method, err)
|
||||
}
|
||||
|
||||
mh := session.receivingMethodHandler()
|
||||
re, _ := jreq.Extra.(*RequestExtra)
|
||||
req := info.newRequest(session, params, re)
|
||||
// mh might be user code, so ensure that it returns the right values for the jsonrpc2 protocol.
|
||||
res, err := mh(ctx, jreq.Method, req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return res, nil
|
||||
}
|
||||
|
||||
// checkRequest checks the given request against the provided method info, to
|
||||
// ensure it is a valid MCP request.
|
||||
//
|
||||
// If valid, the relevant method info is returned. Otherwise, a non-nil error
|
||||
// is returned describing why the request is invalid.
|
||||
//
|
||||
// This is extracted from request handling so that it can be called in the
|
||||
// transport layer to preemptively reject bad requests.
|
||||
func checkRequest(req *jsonrpc.Request, infos map[string]methodInfo) (methodInfo, error) {
|
||||
info, ok := infos[req.Method]
|
||||
if !ok {
|
||||
return methodInfo{}, fmt.Errorf("%w: %q unsupported", jsonrpc2.ErrNotHandled, req.Method)
|
||||
}
|
||||
if info.flags¬ification != 0 && req.IsCall() {
|
||||
return methodInfo{}, fmt.Errorf("%w: unexpected id for %q", jsonrpc2.ErrInvalidRequest, req.Method)
|
||||
}
|
||||
if info.flags¬ification == 0 && !req.IsCall() {
|
||||
return methodInfo{}, fmt.Errorf("%w: missing id for %q", jsonrpc2.ErrInvalidRequest, req.Method)
|
||||
}
|
||||
// missingParamsOK is checked here to catch the common case where "params" is
|
||||
// missing entirely.
|
||||
//
|
||||
// However, it's checked again after unmarshalling to catch the rare but
|
||||
// possible case where "params" is JSON null (see https://go.dev/issue/33835).
|
||||
if info.flags&missingParamsOK == 0 && len(req.Params) == 0 {
|
||||
return methodInfo{}, fmt.Errorf("%w: missing required \"params\"", jsonrpc2.ErrInvalidRequest)
|
||||
}
|
||||
return info, nil
|
||||
}
|
||||
|
||||
// methodInfo is information about sending and receiving a method.
|
||||
type methodInfo struct {
|
||||
// flags is a collection of flags controlling how the JSONRPC method is
|
||||
// handled. See individual flag values for documentation.
|
||||
flags methodFlags
|
||||
// Unmarshal params from the wire into a Params struct.
|
||||
// Used on the receive side.
|
||||
unmarshalParams func(json.RawMessage) (Params, error)
|
||||
newRequest func(Session, Params, *RequestExtra) Request
|
||||
// Run the code when a call to the method is received.
|
||||
// Used on the receive side.
|
||||
handleMethod MethodHandler
|
||||
// Create a pointer to a Result struct.
|
||||
// Used on the send side.
|
||||
newResult func() Result
|
||||
}
|
||||
|
||||
// The following definitions support converting from typed to untyped method handlers.
|
||||
// Type parameter meanings:
|
||||
// - S: sessions
|
||||
// - P: params
|
||||
// - R: results
|
||||
|
||||
// A typedMethodHandler is like a MethodHandler, but with type information.
|
||||
type (
|
||||
typedClientMethodHandler[P Params, R Result] func(context.Context, *ClientRequest[P]) (R, error)
|
||||
typedServerMethodHandler[P Params, R Result] func(context.Context, *ServerRequest[P]) (R, error)
|
||||
)
|
||||
|
||||
type paramsPtr[T any] interface {
|
||||
*T
|
||||
Params
|
||||
}
|
||||
|
||||
type methodFlags int
|
||||
|
||||
const (
|
||||
notification methodFlags = 1 << iota // method is a notification, not request
|
||||
missingParamsOK // params may be missing or null
|
||||
)
|
||||
|
||||
func newClientMethodInfo[P paramsPtr[T], R Result, T any](d typedClientMethodHandler[P, R], flags methodFlags) methodInfo {
|
||||
mi := newMethodInfo[P, R](flags)
|
||||
mi.newRequest = func(s Session, p Params, _ *RequestExtra) Request {
|
||||
r := &ClientRequest[P]{Session: s.(*ClientSession)}
|
||||
if p != nil {
|
||||
r.Params = p.(P)
|
||||
}
|
||||
return r
|
||||
}
|
||||
mi.handleMethod = MethodHandler(func(ctx context.Context, _ string, req Request) (Result, error) {
|
||||
return d(ctx, req.(*ClientRequest[P]))
|
||||
})
|
||||
return mi
|
||||
}
|
||||
|
||||
func newServerMethodInfo[P paramsPtr[T], R Result, T any](d typedServerMethodHandler[P, R], flags methodFlags) methodInfo {
|
||||
mi := newMethodInfo[P, R](flags)
|
||||
mi.newRequest = func(s Session, p Params, re *RequestExtra) Request {
|
||||
r := &ServerRequest[P]{Session: s.(*ServerSession), Extra: re}
|
||||
if p != nil {
|
||||
r.Params = p.(P)
|
||||
}
|
||||
return r
|
||||
}
|
||||
mi.handleMethod = MethodHandler(func(ctx context.Context, _ string, req Request) (Result, error) {
|
||||
return d(ctx, req.(*ServerRequest[P]))
|
||||
})
|
||||
return mi
|
||||
}
|
||||
|
||||
// newMethodInfo creates a methodInfo from a typedMethodHandler.
|
||||
//
|
||||
// If isRequest is set, the method is treated as a request rather than a
|
||||
// notification.
|
||||
func newMethodInfo[P paramsPtr[T], R Result, T any](flags methodFlags) methodInfo {
|
||||
return methodInfo{
|
||||
flags: flags,
|
||||
unmarshalParams: func(m json.RawMessage) (Params, error) {
|
||||
var p P
|
||||
if m != nil {
|
||||
if err := internaljson.Unmarshal(m, &p); err != nil {
|
||||
return nil, fmt.Errorf("unmarshaling %q into a %T: %w", m, p, err)
|
||||
}
|
||||
}
|
||||
// We must check missingParamsOK here, in addition to checkRequest, to
|
||||
// catch the edge cases where "params" is set to JSON null.
|
||||
// See also https://go.dev/issue/33835.
|
||||
//
|
||||
// We need to ensure that p is non-null to guard against crashes, as our
|
||||
// internal code or externally provided handlers may assume that params
|
||||
// is non-null.
|
||||
if flags&missingParamsOK == 0 && p == nil {
|
||||
return nil, fmt.Errorf("%w: missing required \"params\"", jsonrpc2.ErrInvalidRequest)
|
||||
}
|
||||
return orZero[Params](p), nil
|
||||
},
|
||||
// newResult is used on the send side, to construct the value to unmarshal the result into.
|
||||
// R is a pointer to a result struct. There is no way to "unpointer" it without reflection.
|
||||
// TODO(jba): explore generic approaches to this, perhaps by treating R in
|
||||
// the signature as the unpointered type.
|
||||
newResult: func() Result { return reflect.New(reflect.TypeFor[R]().Elem()).Interface().(R) },
|
||||
}
|
||||
}
|
||||
|
||||
// serverMethod is glue for creating a typedMethodHandler from a method on Server.
|
||||
func serverMethod[P Params, R Result](
|
||||
f func(*Server, context.Context, *ServerRequest[P]) (R, error),
|
||||
) typedServerMethodHandler[P, R] {
|
||||
return func(ctx context.Context, req *ServerRequest[P]) (R, error) {
|
||||
return f(req.Session.server, ctx, req)
|
||||
}
|
||||
}
|
||||
|
||||
// clientMethod is glue for creating a typedMethodHandler from a method on Client.
|
||||
func clientMethod[P Params, R Result](
|
||||
f func(*Client, context.Context, *ClientRequest[P]) (R, error),
|
||||
) typedClientMethodHandler[P, R] {
|
||||
return func(ctx context.Context, req *ClientRequest[P]) (R, error) {
|
||||
return f(req.Session.client, ctx, req)
|
||||
}
|
||||
}
|
||||
|
||||
// serverSessionMethod is glue for creating a typedServerMethodHandler from a method on ServerSession.
|
||||
func serverSessionMethod[P Params, R Result](f func(*ServerSession, context.Context, P) (R, error)) typedServerMethodHandler[P, R] {
|
||||
return func(ctx context.Context, req *ServerRequest[P]) (R, error) {
|
||||
return f(req.GetSession().(*ServerSession), ctx, req.Params)
|
||||
}
|
||||
}
|
||||
|
||||
// clientSessionMethod is glue for creating a typedMethodHandler from a method on ServerSession.
|
||||
func clientSessionMethod[P Params, R Result](f func(*ClientSession, context.Context, P) (R, error)) typedClientMethodHandler[P, R] {
|
||||
return func(ctx context.Context, req *ClientRequest[P]) (R, error) {
|
||||
return f(req.GetSession().(*ClientSession), ctx, req.Params)
|
||||
}
|
||||
}
|
||||
|
||||
// MCP-specific error codes.
|
||||
const (
|
||||
// CodeResourceNotFound indicates that a requested resource could not be found.
|
||||
CodeResourceNotFound = -32002
|
||||
// CodeURLElicitationRequired indicates that the server requires URL elicitation
|
||||
// before processing the request. The client should execute the elicitation handler
|
||||
// with the elicitations provided in the error data.
|
||||
CodeURLElicitationRequired = -32042
|
||||
)
|
||||
|
||||
// URLElicitationRequiredError returns an error indicating that URL elicitation is required
|
||||
// before the request can be processed. The elicitations parameter should contain the
|
||||
// elicitation requests that must be completed.
|
||||
func URLElicitationRequiredError(elicitations []*ElicitParams) error {
|
||||
// Validate that all elicitations are URL mode
|
||||
for _, elicit := range elicitations {
|
||||
mode := elicit.Mode
|
||||
if mode == "" {
|
||||
mode = "form" // default mode
|
||||
}
|
||||
if mode != "url" {
|
||||
panic(fmt.Sprintf("URLElicitationRequiredError requires all elicitations to be URL mode, got %q", mode))
|
||||
}
|
||||
}
|
||||
|
||||
data, err := json.Marshal(map[string]any{
|
||||
"elicitations": elicitations,
|
||||
})
|
||||
if err != nil {
|
||||
// This should never happen with valid ElicitParams
|
||||
panic(fmt.Sprintf("failed to marshal elicitations: %v", err))
|
||||
}
|
||||
return &jsonrpc.Error{
|
||||
Code: CodeURLElicitationRequired,
|
||||
Message: "URL elicitation required",
|
||||
Data: json.RawMessage(data),
|
||||
}
|
||||
}
|
||||
|
||||
// Internal error codes
|
||||
const (
|
||||
// The error code if the method exists and was called properly, but the peer does not support it.
|
||||
//
|
||||
// TODO(rfindley): this code is wrong, and we should fix it to be
|
||||
// consistent with other SDKs.
|
||||
codeUnsupportedMethod = -31001
|
||||
)
|
||||
|
||||
// notifySessions calls Notify on all the sessions.
|
||||
// Should be called on a copy of the peer sessions.
|
||||
// The logger must be non-nil.
|
||||
func notifySessions[S Session, P Params](sessions []S, method string, params P, logger *slog.Logger) {
|
||||
if sessions == nil {
|
||||
return
|
||||
}
|
||||
// Notify with the background context, so the messages are sent on the
|
||||
// standalone stream.
|
||||
// TODO: make this timeout configurable, or call handleNotify asynchronously.
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// TODO: there's a potential spec violation here, when the feature list
|
||||
// changes before the session (client or server) is initialized.
|
||||
for _, s := range sessions {
|
||||
req := newRequest(s, params)
|
||||
if err := handleNotify(ctx, method, req); err != nil {
|
||||
logger.Warn(fmt.Sprintf("calling %s: %v", method, err))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func newRequest[S Session, P Params](s S, p P) Request {
|
||||
switch s := any(s).(type) {
|
||||
case *ClientSession:
|
||||
return &ClientRequest[P]{Session: s, Params: p}
|
||||
case *ServerSession:
|
||||
return &ServerRequest[P]{Session: s, Params: p}
|
||||
default:
|
||||
panic("bad session")
|
||||
}
|
||||
}
|
||||
|
||||
// Meta is additional metadata for requests, responses and other types.
|
||||
type Meta map[string]any
|
||||
|
||||
// GetMeta returns metadata from a value.
|
||||
func (m Meta) GetMeta() map[string]any { return m }
|
||||
|
||||
// SetMeta sets the metadata on a value.
|
||||
func (m *Meta) SetMeta(x map[string]any) { *m = x }
|
||||
|
||||
const progressTokenKey = "progressToken"
|
||||
|
||||
func getProgressToken(p Params) any {
|
||||
return p.GetMeta()[progressTokenKey]
|
||||
}
|
||||
|
||||
func setProgressToken(p Params, pt any) {
|
||||
switch pt.(type) {
|
||||
// Support int32 and int64 for atomic.IntNN.
|
||||
case int, int32, int64, string:
|
||||
default:
|
||||
panic(fmt.Sprintf("progress token %v is of type %[1]T, not int or string", pt))
|
||||
}
|
||||
m := p.GetMeta()
|
||||
if m == nil {
|
||||
m = map[string]any{}
|
||||
}
|
||||
m[progressTokenKey] = pt
|
||||
}
|
||||
|
||||
// A Request is a method request with parameters and additional information, such as the session.
|
||||
// Request is implemented by [*ClientRequest] and [*ServerRequest].
|
||||
type Request interface {
|
||||
isRequest()
|
||||
GetSession() Session
|
||||
GetParams() Params
|
||||
// GetExtra returns the Extra field for ServerRequests, and nil for ClientRequests.
|
||||
GetExtra() *RequestExtra
|
||||
}
|
||||
|
||||
// A ClientRequest is a request to a client.
|
||||
type ClientRequest[P Params] struct {
|
||||
Session *ClientSession
|
||||
Params P
|
||||
}
|
||||
|
||||
// A ServerRequest is a request to a server.
|
||||
type ServerRequest[P Params] struct {
|
||||
Session *ServerSession
|
||||
Params P
|
||||
Extra *RequestExtra
|
||||
}
|
||||
|
||||
// RequestExtra is extra information included in requests, typically from
|
||||
// the transport layer.
|
||||
type RequestExtra struct {
|
||||
TokenInfo *auth.TokenInfo // bearer token info (e.g. from OAuth) if any
|
||||
Header http.Header // header from HTTP request, if any
|
||||
|
||||
// If set, CloseSSEStream explicitly closes the current SSE request stream.
|
||||
//
|
||||
// [SEP-1699] introduced server-side SSE stream disconnection: for
|
||||
// long-running requests, servers may opt to close the SSE stream and
|
||||
// ask the client to retry at a later time. CloseSSEStream implements this
|
||||
// feature; if RetryAfter is set, an event is sent with a `retry:` field
|
||||
// to configure the reconnection delay.
|
||||
//
|
||||
// [SEP-1699]: https://github.com/modelcontextprotocol/modelcontextprotocol/issues/1699
|
||||
CloseSSEStream func(CloseSSEStreamArgs)
|
||||
}
|
||||
|
||||
// CloseSSEStreamArgs are arguments for [RequestExtra.CloseSSEStream].
|
||||
type CloseSSEStreamArgs struct {
|
||||
// RetryAfter configures the reconnection delay sent to the client via the
|
||||
// SSE retry field. If zero, no retry field is sent.
|
||||
RetryAfter time.Duration
|
||||
}
|
||||
|
||||
func (*ClientRequest[P]) isRequest() {}
|
||||
func (*ServerRequest[P]) isRequest() {}
|
||||
|
||||
func (r *ClientRequest[P]) GetSession() Session { return r.Session }
|
||||
func (r *ServerRequest[P]) GetSession() Session { return r.Session }
|
||||
|
||||
func (r *ClientRequest[P]) GetParams() Params { return r.Params }
|
||||
func (r *ServerRequest[P]) GetParams() Params { return r.Params }
|
||||
|
||||
func (r *ClientRequest[P]) GetExtra() *RequestExtra { return nil }
|
||||
func (r *ServerRequest[P]) GetExtra() *RequestExtra { return r.Extra }
|
||||
|
||||
func serverRequestFor[P Params](s *ServerSession, p P) *ServerRequest[P] {
|
||||
return &ServerRequest[P]{Session: s, Params: p}
|
||||
}
|
||||
|
||||
func clientRequestFor[P Params](s *ClientSession, p P) *ClientRequest[P] {
|
||||
return &ClientRequest[P]{Session: s, Params: p}
|
||||
}
|
||||
|
||||
// Params is a parameter (input) type for an MCP call or notification.
|
||||
type Params interface {
|
||||
// GetMeta returns metadata from a value.
|
||||
GetMeta() map[string]any
|
||||
// SetMeta sets the metadata on a value.
|
||||
SetMeta(map[string]any)
|
||||
|
||||
// isParams discourages implementation of Params outside of this package.
|
||||
isParams()
|
||||
}
|
||||
|
||||
// RequestParams is a parameter (input) type for an MCP request.
|
||||
type RequestParams interface {
|
||||
Params
|
||||
|
||||
// GetProgressToken returns the progress token from the params' Meta field, or nil
|
||||
// if there is none.
|
||||
GetProgressToken() any
|
||||
|
||||
// SetProgressToken sets the given progress token into the params' Meta field.
|
||||
// It panics if its argument is not an int or a string.
|
||||
SetProgressToken(any)
|
||||
}
|
||||
|
||||
// Result is a result of an MCP call.
|
||||
type Result interface {
|
||||
// isResult discourages implementation of Result outside of this package.
|
||||
isResult()
|
||||
|
||||
// GetMeta returns metadata from a value.
|
||||
GetMeta() map[string]any
|
||||
// SetMeta sets the metadata on a value.
|
||||
SetMeta(map[string]any)
|
||||
}
|
||||
|
||||
// emptyResult is returned by methods that have no result, like ping.
|
||||
// Those methods cannot return nil, because jsonrpc2 cannot handle nils.
|
||||
type emptyResult struct{}
|
||||
|
||||
func (*emptyResult) isResult() {}
|
||||
func (*emptyResult) GetMeta() map[string]any { panic("should never be called") }
|
||||
func (*emptyResult) SetMeta(map[string]any) { panic("should never be called") }
|
||||
|
||||
type listParams interface {
|
||||
// Returns a pointer to the param's Cursor field.
|
||||
cursorPtr() *string
|
||||
}
|
||||
|
||||
type listResult[T any] interface {
|
||||
// Returns a pointer to the param's NextCursor field.
|
||||
nextCursorPtr() *string
|
||||
}
|
||||
|
||||
// keepaliveSession represents a session that supports keepalive functionality.
|
||||
type keepaliveSession interface {
|
||||
Ping(ctx context.Context, params *PingParams) error
|
||||
Close() error
|
||||
}
|
||||
|
||||
// startKeepalive starts the keepalive mechanism for a session.
|
||||
// It assigns the cancel function to the provided cancelPtr and starts a goroutine
|
||||
// that sends ping messages at the specified interval.
|
||||
func startKeepalive(session keepaliveSession, interval time.Duration, cancelPtr *context.CancelFunc) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
// Assign cancel function before starting goroutine to avoid race condition.
|
||||
// We cannot return it because the caller may need to cancel during the
|
||||
// window between goroutine scheduling and function return.
|
||||
*cancelPtr = cancel
|
||||
|
||||
go func() {
|
||||
ticker := time.NewTicker(interval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
pingCtx, pingCancel := context.WithTimeout(context.Background(), interval/2)
|
||||
err := session.Ping(pingCtx, nil)
|
||||
pingCancel()
|
||||
if err != nil {
|
||||
// Ping failed, close the session
|
||||
_ = session.Close()
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
+489
@@ -0,0 +1,489 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sync"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/internal/jsonrpc2"
|
||||
"github.com/modelcontextprotocol/go-sdk/jsonrpc"
|
||||
)
|
||||
|
||||
// This file implements support for SSE (HTTP with server-sent events)
|
||||
// transport server and client.
|
||||
// https://modelcontextprotocol.io/specification/2024-11-05/basic/transports
|
||||
//
|
||||
// The transport is simple, at least relative to the new streamable transport
|
||||
// introduced in the 2025-03-26 version of the spec. In short:
|
||||
//
|
||||
// 1. Sessions are initiated via a hanging GET request, which streams
|
||||
// server->client messages as SSE 'message' events.
|
||||
// 2. The first event in the SSE stream must be an 'endpoint' event that
|
||||
// informs the client of the session endpoint.
|
||||
// 3. The client POSTs client->server messages to the session endpoint.
|
||||
//
|
||||
// Therefore, the each new GET request hands off its responsewriter to an
|
||||
// [SSEServerTransport] type that abstracts the transport as follows:
|
||||
// - Write writes a new event to the responseWriter, or fails if the GET has
|
||||
// exited.
|
||||
// - Read reads off a message queue that is pushed to via POST requests.
|
||||
// - Close causes the hanging GET to exit.
|
||||
|
||||
// SSEHandler is an http.Handler that serves SSE-based MCP sessions as defined by
|
||||
// the [2024-11-05 version] of the MCP spec.
|
||||
//
|
||||
// [2024-11-05 version]: https://modelcontextprotocol.io/specification/2024-11-05/basic/transports
|
||||
type SSEHandler struct {
|
||||
getServer func(request *http.Request) *Server
|
||||
opts SSEOptions
|
||||
onConnection func(*ServerSession) // for testing; must not block
|
||||
|
||||
mu sync.Mutex
|
||||
sessions map[string]*SSEServerTransport
|
||||
}
|
||||
|
||||
// SSEOptions specifies options for an [SSEHandler].
|
||||
// for now, it is empty, but may be extended in future.
|
||||
// https://github.com/modelcontextprotocol/go-sdk/issues/507
|
||||
type SSEOptions struct{}
|
||||
|
||||
// NewSSEHandler returns a new [SSEHandler] that creates and manages MCP
|
||||
// sessions created via incoming HTTP requests.
|
||||
//
|
||||
// Sessions are created when the client issues a GET request to the server,
|
||||
// which must accept text/event-stream responses (server-sent events).
|
||||
// For each such request, a new [SSEServerTransport] is created with a distinct
|
||||
// messages endpoint, and connected to the server returned by getServer.
|
||||
// The SSEHandler also handles requests to the message endpoints, by
|
||||
// delegating them to the relevant server transport.
|
||||
//
|
||||
// The getServer function may return a distinct [Server] for each new
|
||||
// request, or reuse an existing server. If it returns nil, the handler
|
||||
// will return a 400 Bad Request.
|
||||
func NewSSEHandler(getServer func(request *http.Request) *Server, opts *SSEOptions) *SSEHandler {
|
||||
s := &SSEHandler{
|
||||
getServer: getServer,
|
||||
sessions: make(map[string]*SSEServerTransport),
|
||||
}
|
||||
|
||||
if opts != nil {
|
||||
s.opts = *opts
|
||||
}
|
||||
|
||||
return s
|
||||
}
|
||||
|
||||
// A SSEServerTransport is a logical SSE session created through a hanging GET
|
||||
// request.
|
||||
//
|
||||
// Use [SSEServerTransport.Connect] to initiate the flow of messages.
|
||||
//
|
||||
// When connected, it returns the following [Connection] implementation:
|
||||
// - Writes are SSE 'message' events to the GET response.
|
||||
// - Reads are received from POSTs to the session endpoint, via
|
||||
// [SSEServerTransport.ServeHTTP].
|
||||
// - Close terminates the hanging GET.
|
||||
//
|
||||
// The transport is itself an [http.Handler]. It is the caller's responsibility
|
||||
// to ensure that the resulting transport serves HTTP requests on the given
|
||||
// session endpoint.
|
||||
//
|
||||
// Each SSEServerTransport may be connected (via [Server.Connect]) at most
|
||||
// once, since [SSEServerTransport.ServeHTTP] serves messages to the connected
|
||||
// session.
|
||||
//
|
||||
// Most callers should instead use an [SSEHandler], which transparently handles
|
||||
// the delegation to SSEServerTransports.
|
||||
type SSEServerTransport struct {
|
||||
// Endpoint is the endpoint for this session, where the client can POST
|
||||
// messages.
|
||||
Endpoint string
|
||||
|
||||
// Response is the hanging response body to the incoming GET request.
|
||||
Response http.ResponseWriter
|
||||
|
||||
// incoming is the queue of incoming messages.
|
||||
// It is never closed, and by convention, incoming is non-nil if and only if
|
||||
// the transport is connected.
|
||||
incoming chan jsonrpc.Message
|
||||
|
||||
// We must guard both pushes to the incoming queue and writes to the response
|
||||
// writer, because incoming POST requests are arbitrarily concurrent and we
|
||||
// need to ensure we don't write push to the queue, or write to the
|
||||
// ResponseWriter, after the session GET request exits.
|
||||
mu sync.Mutex // also guards writes to Response
|
||||
closed bool // set when the stream is closed
|
||||
done chan struct{} // closed when the connection is closed
|
||||
}
|
||||
|
||||
// ServeHTTP handles POST requests to the transport endpoint.
|
||||
func (t *SSEServerTransport) ServeHTTP(w http.ResponseWriter, req *http.Request) {
|
||||
if t.incoming == nil {
|
||||
http.Error(w, "session not connected", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
// Read and parse the message.
|
||||
data, err := io.ReadAll(req.Body)
|
||||
if err != nil {
|
||||
http.Error(w, "failed to read body", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
// Optionally, we could just push the data onto a channel, and let the
|
||||
// message fail to parse when it is read. This failure seems a bit more
|
||||
// useful
|
||||
msg, err := jsonrpc2.DecodeMessage(data)
|
||||
if err != nil {
|
||||
http.Error(w, "failed to parse body", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
if req, ok := msg.(*jsonrpc.Request); ok {
|
||||
if _, err := checkRequest(req, serverMethodInfos); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
}
|
||||
select {
|
||||
case t.incoming <- msg:
|
||||
w.WriteHeader(http.StatusAccepted)
|
||||
case <-t.done:
|
||||
http.Error(w, "session closed", http.StatusBadRequest)
|
||||
}
|
||||
}
|
||||
|
||||
// Connect sends the 'endpoint' event to the client.
|
||||
// See [SSEServerTransport] for more details on the [Connection] implementation.
|
||||
func (t *SSEServerTransport) Connect(context.Context) (Connection, error) {
|
||||
if t.incoming != nil {
|
||||
return nil, fmt.Errorf("already connected")
|
||||
}
|
||||
t.incoming = make(chan jsonrpc.Message, 100)
|
||||
t.done = make(chan struct{})
|
||||
_, err := writeEvent(t.Response, Event{
|
||||
Name: "endpoint",
|
||||
Data: []byte(t.Endpoint),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &sseServerConn{t: t}, nil
|
||||
}
|
||||
|
||||
func (h *SSEHandler) ServeHTTP(w http.ResponseWriter, req *http.Request) {
|
||||
sessionID := req.URL.Query().Get("sessionid")
|
||||
|
||||
// TODO: consider checking Content-Type here. For now, we are lax.
|
||||
|
||||
// For POST requests, the message body is a message to send to a session.
|
||||
if req.Method == http.MethodPost {
|
||||
// Look up the session.
|
||||
if sessionID == "" {
|
||||
http.Error(w, "sessionid must be provided", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
h.mu.Lock()
|
||||
session := h.sessions[sessionID]
|
||||
h.mu.Unlock()
|
||||
if session == nil {
|
||||
http.Error(w, "session not found", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
session.ServeHTTP(w, req)
|
||||
return
|
||||
}
|
||||
|
||||
if req.Method != http.MethodGet {
|
||||
w.Header().Set("Allow", "GET, POST")
|
||||
http.Error(w, "Method Not Allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
// GET requests create a new session, and serve messages over SSE.
|
||||
|
||||
// TODO: it's not entirely documented whether we should check Accept here.
|
||||
// Let's again be lax and assume the client will accept SSE.
|
||||
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
w.Header().Set("Cache-Control", "no-cache")
|
||||
w.Header().Set("Connection", "keep-alive")
|
||||
|
||||
sessionID = rand.Text()
|
||||
endpoint, err := req.URL.Parse("?sessionid=" + sessionID)
|
||||
if err != nil {
|
||||
http.Error(w, "internal error: failed to create endpoint", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
transport := &SSEServerTransport{Endpoint: endpoint.RequestURI(), Response: w}
|
||||
|
||||
// The session is terminated when the request exits.
|
||||
h.mu.Lock()
|
||||
h.sessions[sessionID] = transport
|
||||
h.mu.Unlock()
|
||||
defer func() {
|
||||
h.mu.Lock()
|
||||
delete(h.sessions, sessionID)
|
||||
h.mu.Unlock()
|
||||
}()
|
||||
|
||||
server := h.getServer(req)
|
||||
if server == nil {
|
||||
// The getServer argument to NewSSEHandler returned nil.
|
||||
http.Error(w, "no server available", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
ss, err := server.Connect(req.Context(), transport, nil)
|
||||
if err != nil {
|
||||
http.Error(w, "connection failed", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
if h.onConnection != nil {
|
||||
h.onConnection(ss)
|
||||
}
|
||||
defer ss.Close() // close the transport when the GET exits
|
||||
|
||||
select {
|
||||
case <-req.Context().Done():
|
||||
case <-transport.done:
|
||||
}
|
||||
}
|
||||
|
||||
// sseServerConn implements the [Connection] interface for a single [SSEServerTransport].
|
||||
// It hides the Connection interface from the SSEServerTransport API.
|
||||
type sseServerConn struct {
|
||||
t *SSEServerTransport
|
||||
}
|
||||
|
||||
// TODO(jba): get the session ID. (Not urgent because SSE transports have been removed from the spec.)
|
||||
func (s *sseServerConn) SessionID() string { return "" }
|
||||
|
||||
// Read implements jsonrpc2.Reader.
|
||||
func (s *sseServerConn) Read(ctx context.Context) (jsonrpc.Message, error) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case msg := <-s.t.incoming:
|
||||
return msg, nil
|
||||
case <-s.t.done:
|
||||
return nil, io.EOF
|
||||
}
|
||||
}
|
||||
|
||||
// Write implements jsonrpc2.Writer.
|
||||
func (s *sseServerConn) Write(ctx context.Context, msg jsonrpc.Message) error {
|
||||
if ctx.Err() != nil {
|
||||
return ctx.Err()
|
||||
}
|
||||
|
||||
data, err := jsonrpc2.EncodeMessage(msg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
s.t.mu.Lock()
|
||||
defer s.t.mu.Unlock()
|
||||
|
||||
// Note that it is invalid to write to a ResponseWriter after ServeHTTP has
|
||||
// exited, and so we must lock around this write and check isDone, which is
|
||||
// set before the hanging GET exits.
|
||||
if s.t.closed {
|
||||
return io.EOF
|
||||
}
|
||||
|
||||
_, err = writeEvent(s.t.Response, Event{Name: "message", Data: data})
|
||||
return err
|
||||
}
|
||||
|
||||
// Close implements io.Closer, and closes the session.
|
||||
//
|
||||
// It must be safe to call Close more than once, as the close may
|
||||
// asynchronously be initiated by either the server closing its connection, or
|
||||
// by the hanging GET exiting.
|
||||
func (s *sseServerConn) Close() error {
|
||||
s.t.mu.Lock()
|
||||
defer s.t.mu.Unlock()
|
||||
if !s.t.closed {
|
||||
s.t.closed = true
|
||||
close(s.t.done)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// An SSEClientTransport is a [Transport] that can communicate with an MCP
|
||||
// endpoint serving the SSE transport defined by the 2024-11-05 version of the
|
||||
// spec.
|
||||
//
|
||||
// https://modelcontextprotocol.io/specification/2024-11-05/basic/transports
|
||||
type SSEClientTransport struct {
|
||||
// Endpoint is the SSE endpoint to connect to.
|
||||
Endpoint string
|
||||
|
||||
// HTTPClient is the client to use for making HTTP requests. If nil,
|
||||
// http.DefaultClient is used.
|
||||
HTTPClient *http.Client
|
||||
}
|
||||
|
||||
// Connect connects through the client endpoint.
|
||||
func (c *SSEClientTransport) Connect(ctx context.Context) (Connection, error) {
|
||||
parsedURL, err := url.Parse(c.Endpoint)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid endpoint: %v", err)
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", c.Endpoint, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
httpClient := c.HTTPClient
|
||||
if httpClient == nil {
|
||||
httpClient = http.DefaultClient
|
||||
}
|
||||
req.Header.Set("Accept", "text/event-stream")
|
||||
resp, err := httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Check HTTP status code before attempting to parse SSE events.
|
||||
// This ensures proper error reporting for authentication failures (401),
|
||||
// authorization failures (403), and other HTTP errors.
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
resp.Body.Close()
|
||||
return nil, fmt.Errorf("failed to connect: %s", http.StatusText(resp.StatusCode))
|
||||
}
|
||||
|
||||
msgEndpoint, err := func() (*url.URL, error) {
|
||||
var evt Event
|
||||
for evt, err = range scanEvents(resp.Body) {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if evt.Name != "endpoint" {
|
||||
return nil, fmt.Errorf("first event is %q, want %q", evt.Name, "endpoint")
|
||||
}
|
||||
raw := string(evt.Data)
|
||||
return parsedURL.Parse(raw)
|
||||
}()
|
||||
if err != nil {
|
||||
resp.Body.Close()
|
||||
return nil, fmt.Errorf("missing endpoint: %v", err)
|
||||
}
|
||||
|
||||
// From here on, the stream takes ownership of resp.Body.
|
||||
s := &sseClientConn{
|
||||
client: httpClient,
|
||||
msgEndpoint: msgEndpoint,
|
||||
incoming: make(chan []byte, 100),
|
||||
body: resp.Body,
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
|
||||
go func() {
|
||||
defer s.Close() // close the transport when the GET exits
|
||||
|
||||
for evt, err := range scanEvents(resp.Body) {
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case s.incoming <- evt.Data:
|
||||
case <-s.done:
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// An sseClientConn is a logical jsonrpc2 connection that implements the client
|
||||
// half of the SSE protocol:
|
||||
// - Writes are POSTS to the session endpoint.
|
||||
// - Reads are SSE 'message' events, and pushes them onto a buffered channel.
|
||||
// - Close terminates the GET request.
|
||||
type sseClientConn struct {
|
||||
client *http.Client // HTTP client to use for requests
|
||||
msgEndpoint *url.URL // session endpoint for POSTs
|
||||
incoming chan []byte // queue of incoming messages
|
||||
|
||||
mu sync.Mutex
|
||||
body io.ReadCloser // body of the hanging GET
|
||||
closed bool // set when the stream is closed
|
||||
done chan struct{} // closed when the stream is closed
|
||||
}
|
||||
|
||||
// TODO(jba): get the session ID. (Not urgent because SSE transports have been removed from the spec.)
|
||||
func (c *sseClientConn) SessionID() string { return "" }
|
||||
|
||||
func (c *sseClientConn) isDone() bool {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
return c.closed
|
||||
}
|
||||
|
||||
func (c *sseClientConn) Read(ctx context.Context) (jsonrpc.Message, error) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
|
||||
case <-c.done:
|
||||
return nil, io.EOF
|
||||
|
||||
case data := <-c.incoming:
|
||||
// TODO(rfindley): do we really need to check this? We receive from c.done above.
|
||||
if c.isDone() {
|
||||
return nil, io.EOF
|
||||
}
|
||||
msg, err := jsonrpc2.DecodeMessage(data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return msg, nil
|
||||
}
|
||||
}
|
||||
|
||||
func (c *sseClientConn) Write(ctx context.Context, msg jsonrpc.Message) error {
|
||||
data, err := jsonrpc2.EncodeMessage(msg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if c.isDone() {
|
||||
return io.EOF
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", c.msgEndpoint.String(), bytes.NewReader(data))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp, err := c.client.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return fmt.Errorf("failed to write: %s", resp.Status)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *sseClientConn) Close() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if !c.closed {
|
||||
c.closed = true
|
||||
_ = c.body.Close()
|
||||
close(c.done)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
+2221
File diff suppressed because it is too large
Load Diff
+226
@@ -0,0 +1,226 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// TODO: move client-side streamable HTTP logic from streamable.go to this file.
|
||||
|
||||
package mcp
|
||||
|
||||
/*
|
||||
Streamable HTTP Client Design
|
||||
|
||||
This document describes the client-side implementation of the MCP streamable
|
||||
HTTP transport, as defined by the MCP spec:
|
||||
https://modelcontextprotocol.io/specification/2025-11-25/basic/transports#streamable-http
|
||||
|
||||
# Overview
|
||||
|
||||
The client-side streamable transport allows an MCP client to communicate with a
|
||||
server over HTTP, sending messages via POST and receiving responses via either
|
||||
JSON or server-sent events (SSE). The implementation consists of two main
|
||||
components:
|
||||
|
||||
┌─────────────────────────────────────────────────────────────────┐
|
||||
│ [StreamableClientTransport] │
|
||||
│ Transport configuration; creates connections via Connect() │
|
||||
└─────────────────────────────────────────────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌─────────────────────────────────────────────────────────────────┐
|
||||
│ [streamableClientConn] │
|
||||
│ Connection implementation; handles HTTP request/response │
|
||||
└─────────────────────────────────────────────────────────────────┘
|
||||
│
|
||||
├──────────────────────────────────────┐
|
||||
▼ ▼
|
||||
┌─────────────────────────────────────────┐ ┌────────────────────────────────────┐
|
||||
│ POST request handlers │ │ Standalone SSE stream │
|
||||
│ (one per outgoing message/call) │ │ (server-initiated messages) │
|
||||
└─────────────────────────────────────────┘ └────────────────────────────────────┘
|
||||
|
||||
# Sessions
|
||||
|
||||
The client maintains a session with the server, identified by a session ID
|
||||
(Mcp-Session-Id header):
|
||||
|
||||
- Session ID is received from the server after initialization
|
||||
- Client includes the session ID in all subsequent requests
|
||||
- Session ends when the client calls Close() (sends DELETE) or server returns 404
|
||||
|
||||
[streamableClientConn] stores the session state:
|
||||
- [streamableClientConn.sessionID]: Server-assigned session identifier
|
||||
- [streamableClientConn.initializedResult]: Protocol version and server capabilities
|
||||
|
||||
# Connection Lifecycle
|
||||
|
||||
1. Connect: [StreamableClientTransport.Connect] creates a [streamableClientConn]
|
||||
with a detached context for the connection's lifetime. The context is detached
|
||||
to prevent the standalone SSE stream from being cancelled when the original
|
||||
Connect context times out.
|
||||
|
||||
2. Initialize: The MCP client sends initialize/initialized messages. Upon
|
||||
receiving [InitializeResult], the connection:
|
||||
- Stores the negotiated protocol version for the Mcp-Protocol-Version header
|
||||
- Captures the session ID from the Mcp-Session-Id response header
|
||||
- Starts the standalone SSE stream via [streamableClientConn.connectStandaloneSSE]
|
||||
|
||||
3. Operation: Messages are sent via POST, responses received via JSON or SSE.
|
||||
|
||||
4. Close: [streamableClientConn.Close] sends a DELETE request to terminate
|
||||
the session (unless the session is already gone), then cancels the connection
|
||||
context to clean up the standalone SSE stream.
|
||||
|
||||
# Sending Messages (Write)
|
||||
|
||||
[streamableClientConn.Write] sends all outgoing messages via HTTP POST:
|
||||
|
||||
POST /endpoint
|
||||
Content-Type: application/json
|
||||
Accept: application/json, text/event-stream
|
||||
Mcp-Protocol-Version: <negotiated version>
|
||||
Mcp-Session-Id: <session ID, if established>
|
||||
|
||||
<JSON-RPC message>
|
||||
|
||||
The server may respond with:
|
||||
- 202 Accepted: Message received, no response body (notifications/responses)
|
||||
- 200 OK with application/json: Single JSON-RPC response
|
||||
- 200 OK with text/event-stream: SSE stream of responses
|
||||
|
||||
# Receiving Messages (Read)
|
||||
|
||||
[streamableClientConn.Read] returns messages from the [streamableClientConn.incoming]
|
||||
channel, which is populated by multiple concurrent goroutines:
|
||||
|
||||
1. POST response handlers ([streamableClientConn.handleJSON] and
|
||||
[streamableClientConn.handleSSE]): Process responses from POST requests
|
||||
|
||||
2. Standalone SSE stream: Receives server-initiated requests and notifications
|
||||
|
||||
The client handles both response formats:
|
||||
- JSON: [streamableClientConn.handleJSON] reads body, decodes message
|
||||
- SSE: [streamableClientConn.handleSSE] scans events, decodes each message
|
||||
|
||||
# Standalone SSE Stream
|
||||
|
||||
After initialization, [streamableClientConn.sessionUpdated] triggers
|
||||
[streamableClientConn.connectStandaloneSSE] to open a GET request for
|
||||
server-initiated messages:
|
||||
|
||||
GET /endpoint
|
||||
Accept: text/event-stream
|
||||
Mcp-Session-Id: <session ID>
|
||||
|
||||
Stream behavior:
|
||||
- Optional: Server may return 405 Method Not Allowed (spec-compliant) or
|
||||
other 4xx errors (tolerated in non-strict mode for compatibility)
|
||||
- Persistent: Runs for the connection lifetime in a background goroutine
|
||||
- Resumable: Uses Last-Event-ID header on reconnection if server provides event IDs
|
||||
- Reconnects: Automatic reconnection with exponential backoff on interruption
|
||||
|
||||
# Stream Resumption
|
||||
|
||||
When an SSE stream (standalone or POST response) is interrupted, the client
|
||||
attempts to reconnect using [streamableClientConn.connectSSE]:
|
||||
|
||||
Event ID tracking:
|
||||
- [streamableClientConn.processStream] tracks the last received event ID
|
||||
- On reconnection, the Last-Event-ID header is set to resume from that point
|
||||
- Server replays missed events if it has an [EventStore] configured
|
||||
|
||||
See [calculateReconnectDelay] for the reconnect delay details.
|
||||
|
||||
Server-initiated reconnection (SEP-1699)
|
||||
- SSE retry field: Sets the delay for the next reconnect attempt
|
||||
- If server doesn't provide event IDs, non-standalone streams don't reconnect
|
||||
|
||||
# Response Formats
|
||||
|
||||
The client must handle two response formats from POST requests:
|
||||
|
||||
1. application/json: Single JSON-RPC response
|
||||
- Body contains one JSON-RPC message
|
||||
- Handled by [streamableClientConn.handleJSON]
|
||||
- Simpler but doesn't support streaming or server-initiated messages
|
||||
|
||||
2. text/event-stream: SSE stream of messages
|
||||
- Body contains SSE events with JSON-RPC messages
|
||||
- Handled by [streamableClientConn.handleSSE]
|
||||
- Supports multiple messages and server-initiated communication
|
||||
- Stream completes when the response to the originating call is received
|
||||
|
||||
# HTTP Methods
|
||||
|
||||
- POST: Send JSON-RPC messages (requests, responses, notifications)
|
||||
- Used by [streamableClientConn.Write]
|
||||
- Response may be JSON or SSE
|
||||
|
||||
- GET: Open or resume SSE stream for server-initiated messages
|
||||
- Used by [streamableClientConn.connectSSE]
|
||||
- Always expects text/event-stream response (or 405)
|
||||
|
||||
- DELETE: Terminate the session
|
||||
- Used by [streamableClientConn.Close]
|
||||
- Skipped if session is already known to be gone ([ErrSessionMissing])
|
||||
|
||||
# Error Handling
|
||||
|
||||
Errors are categorized and handled differently:
|
||||
|
||||
1. Transient (recoverable via reconnection):
|
||||
- Network interruption during SSE streaming
|
||||
- Connection reset or timeout
|
||||
- Triggers reconnection in [streamableClientConn.handleSSE]
|
||||
|
||||
2. Terminal (breaks the connection):
|
||||
- 404 Not Found: Session terminated by server ([ErrSessionMissing])
|
||||
- Message decode errors: Protocol violation
|
||||
- Context cancellation: Client closed connection
|
||||
- Mismatched session IDs: Protocol error
|
||||
- See issue #683: our terminal errors are too strict.
|
||||
|
||||
Terminal errors are stored via [streamableClientConn.fail] and returned by
|
||||
subsequent [streamableClientConn.Read] calls. The [streamableClientConn.failed]
|
||||
channel signals that the connection is broken.
|
||||
|
||||
Special case: [ErrSessionMissing] indicates the server has terminated the session,
|
||||
so [streamableClientConn.Close] skips the DELETE request.
|
||||
|
||||
# Protocol Version Header
|
||||
|
||||
After initialization, all requests include:
|
||||
|
||||
Mcp-Protocol-Version: <negotiated version>
|
||||
|
||||
This header (set by [streamableClientConn.setMCPHeaders]):
|
||||
- Allows the server to handle requests per the negotiated protocol
|
||||
- Is omitted before initialization completes
|
||||
- Uses the version from [streamableClientConn.initializedResult]
|
||||
|
||||
# Key Implementation Details
|
||||
|
||||
[StreamableClientTransport] configuration:
|
||||
- [StreamableClientTransport.Endpoint]: URL of the MCP server
|
||||
- [StreamableClientTransport.HTTPClient]: Custom HTTP client (optional)
|
||||
- [StreamableClientTransport.MaxRetries]: Reconnection attempts (default 5)
|
||||
|
||||
[streamableClientConn] handles the [Connection] interface:
|
||||
- [streamableClientConn.Read]: Returns messages from incoming channel
|
||||
- [streamableClientConn.Write]: Sends messages via POST, starts response handlers
|
||||
- [streamableClientConn.Close]: Sends DELETE, cancels context, closes done channel
|
||||
|
||||
State management:
|
||||
- [streamableClientConn.incoming]: Buffered channel for received messages
|
||||
- [streamableClientConn.sessionID]: Server-assigned session identifier
|
||||
- [streamableClientConn.initializedResult]: Cached for protocol version header
|
||||
- [streamableClientConn.failed]: Channel closed on terminal error
|
||||
- [streamableClientConn.done]: Channel closed on graceful shutdown
|
||||
- [streamableClientConn.ctx]: Detached context for connection lifetime
|
||||
- [streamableClientConn.cancel]: Cancels ctx to terminate SSE streams
|
||||
|
||||
Context handling:
|
||||
- Connection context is detached from [StreamableClientTransport.Connect] context
|
||||
using [xcontext.Detach] to preserve context values (for auth middleware) while
|
||||
preventing premature cancellation of the standalone SSE stream
|
||||
- Individual POST requests use caller-provided contexts for cancellation
|
||||
*/
|
||||
+160
@@ -0,0 +1,160 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// TODO: move server-side streamable HTTP logic from streamable.go to this file.
|
||||
|
||||
package mcp
|
||||
|
||||
/*
|
||||
Streamable HTTP Server Design
|
||||
|
||||
This document describes the server-side implementation of the MCP streamable
|
||||
HTTP transport, as defined by the MCP spec:
|
||||
https://modelcontextprotocol.io/specification/2025-11-25/basic/transports#streamable-http
|
||||
|
||||
# Overview
|
||||
|
||||
The streamable HTTP transport enables MCP communication over HTTP, with
|
||||
server-sent events (SSE) for server-to-client messages. The implementation
|
||||
consists of several layered components:
|
||||
|
||||
┌─────────────────────────────────────────────────────────────────┐
|
||||
│ [StreamableHTTPHandler] │
|
||||
│ http.Handler that manages sessions and routes HTTP requests │
|
||||
└─────────────────────────────────────────────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌─────────────────────────────────────────────────────────────────┐
|
||||
│ [StreamableServerTransport] │
|
||||
│ transport implementation, one per session; exposes ServeHTTP │
|
||||
└─────────────────────────────────────────────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌─────────────────────────────────────────────────────────────────┐
|
||||
│ [streamableServerConn] │
|
||||
│ Connection implementation, handles message routing │
|
||||
└─────────────────────────────────────────────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌─────────────────────────────────────────────────────────────────┐
|
||||
│ [stream] │
|
||||
│ Logical message channel within a session, may be resumed │
|
||||
└─────────────────────────────────────────────────────────────────┘
|
||||
|
||||
# Sessions
|
||||
|
||||
As with other transports, a session represents a logical MCP connection between
|
||||
a client and server. In the streamable transport, sessions are identified by a
|
||||
unique session ID (Mcp-Session-Id header) and persist across multiple HTTP
|
||||
requests.
|
||||
|
||||
[StreamableHTTPHandler] maintains a map of active sessions ([sessionInfo]),
|
||||
each containing:
|
||||
- The [ServerSession] (MCP-level session state)
|
||||
- The [StreamableServerTransport] (for message I/O)
|
||||
- Optional timeout management for idle session cleanup
|
||||
|
||||
Sessions are created on the first POST request (typically containing the
|
||||
initialize request) and destroyed either by:
|
||||
- Client sending a DELETE request
|
||||
- Session timeout due to inactivity
|
||||
- Server explicitly closing the session
|
||||
|
||||
# Streams
|
||||
|
||||
Within a session, there can be multiple concurrent "streams" - logical channels
|
||||
for message delivery. This is distinct from HTTP streams; a single [stream] may
|
||||
span multiple HTTP request/response cycles (via resumption).
|
||||
|
||||
There are two types of streams:
|
||||
|
||||
1. Optional standalone SSE stream (id = ""):
|
||||
- Created when client sends a GET request to the endpoint
|
||||
- Used for server-initiated messages (requests/notifications to client)
|
||||
- Persists for the lifetime of the session
|
||||
- Only one standalone stream per session
|
||||
|
||||
2. Request streams (id = random string):
|
||||
- Created for each POST request containing JSON-RPC calls
|
||||
- Used to route responses back to the originating HTTP request
|
||||
- Completed when all responses have been sent
|
||||
- Can be resumed via GET with Last-Event-ID if interrupted
|
||||
|
||||
# Message Routing
|
||||
|
||||
When the server writes a message, it must be routed to the correct [stream]:
|
||||
|
||||
- Responses: Routed to the stream that originated the request
|
||||
- Requests/Notifications made during request handling: Routed to the same
|
||||
stream as the triggering request (via context)
|
||||
- Requests/Notifications made outside request handling: Routed to the
|
||||
standalone SSE stream
|
||||
|
||||
This routing is implemented using:
|
||||
- [streamableServerConn.requestStreams] maps request IDs to stream IDs
|
||||
- [idContextKey] is used to store the originating request ID in Context
|
||||
- [streamableServerConn.streams] maps stream IDs to [stream] objects
|
||||
|
||||
# Stream Resumption
|
||||
|
||||
If an HTTP connection is interrupted (network issues, etc.), clients can
|
||||
resume a stream by sending a GET request with the Last-Event-ID header.
|
||||
This requires an [EventStore] to be configured on the server.
|
||||
|
||||
- [EventStore.Open] is called when a new stream is created
|
||||
- [EventStore.Append] is called for each message written to the stream
|
||||
- [EventStore.After] is called to replay messages after a given index
|
||||
- [EventStore.SessionClosed] is called when the session ends
|
||||
|
||||
Event IDs are formatted as "<streamID>_<index>" to identify both the
|
||||
stream and position within that stream (see [formatEventID] and [parseEventID]).
|
||||
|
||||
# Stateless Mode
|
||||
|
||||
For simpler deployments, the handler supports "stateless" mode
|
||||
([StreamableHTTPOptions.Stateless]) where:
|
||||
- No session ID validation is performed
|
||||
- Each request creates a temporary session that's closed after the request
|
||||
- Server-to-client requests are not supported (no way to receive response)
|
||||
|
||||
This mode is useful for simple tool servers that don't need bidirectional
|
||||
communication.
|
||||
|
||||
# Response Formats
|
||||
|
||||
The server can respond to POST requests in two formats:
|
||||
|
||||
1. text/event-stream (default): Messages sent as SSE events, supports
|
||||
streaming multiple messages and server-initiated communication during
|
||||
request handling.
|
||||
|
||||
2. application/json ([StreamableHTTPOptions.JSONResponse]): Single JSON
|
||||
response, simpler but doesn't support streaming. Server-initiated messages
|
||||
during request handling go to the standalone SSE stream instead.
|
||||
|
||||
# HTTP Methods
|
||||
|
||||
- POST: Send JSON-RPC messages (requests, responses, notifications)
|
||||
- GET: Open standalone SSE stream or resume an interrupted stream
|
||||
- DELETE: Terminate the session
|
||||
|
||||
# Key Implementation Details
|
||||
|
||||
The [stream] struct manages delivery of messages to HTTP responses.
|
||||
|
||||
Fields:
|
||||
- [stream.w] is the ResponseWriter for the current HTTP response (non-nil indicates claimed)
|
||||
- [stream.done] is closed to release the hanging HTTP request
|
||||
- [stream.requests] tracks pending request IDs (stream completes when empty)
|
||||
|
||||
Methods:
|
||||
- [stream.deliverLocked] delivers a message to the stream
|
||||
- [stream.close] sends a close event and releases the stream
|
||||
- [stream.release] releases the stream from the HTTP request, allowing resumption
|
||||
|
||||
[streamableServerConn] handles the [Connection] interface:
|
||||
- [streamableServerConn.Read] receives messages from the incoming channel (fed by POST handlers)
|
||||
- [streamableServerConn.Write] routes messages to appropriate streams
|
||||
- [streamableServerConn.Close] terminates the session and notifies the [EventStore]
|
||||
*/
|
||||
+140
@@ -0,0 +1,140 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/google/jsonschema-go/jsonschema"
|
||||
internaljson "github.com/modelcontextprotocol/go-sdk/internal/json"
|
||||
)
|
||||
|
||||
// A ToolHandler handles a call to tools/call.
|
||||
//
|
||||
// This is a low-level API, for use with [Server.AddTool]. It does not do any
|
||||
// pre- or post-processing of the request or result: the params contain raw
|
||||
// arguments, no input validation is performed, and the result is returned to
|
||||
// the user as-is, without any validation of the output.
|
||||
//
|
||||
// Most users will write a [ToolHandlerFor] and install it with the generic
|
||||
// [AddTool] function.
|
||||
//
|
||||
// If ToolHandler returns an error, it is treated as a protocol error. By
|
||||
// contrast, [ToolHandlerFor] automatically populates [CallToolResult.IsError]
|
||||
// and [CallToolResult.Content] accordingly.
|
||||
type ToolHandler func(context.Context, *CallToolRequest) (*CallToolResult, error)
|
||||
|
||||
// A ToolHandlerFor handles a call to tools/call with typed arguments and results.
|
||||
//
|
||||
// Use [AddTool] to add a ToolHandlerFor to a server.
|
||||
//
|
||||
// Unlike [ToolHandler], [ToolHandlerFor] provides significant functionality
|
||||
// out of the box, and enforces that the tool conforms to the MCP spec:
|
||||
// - The In type provides a default input schema for the tool, though it may
|
||||
// be overridden in [AddTool].
|
||||
// - The input value is automatically unmarshaled from req.Params.Arguments.
|
||||
// - The input value is automatically validated against its input schema.
|
||||
// Invalid input is rejected before getting to the handler.
|
||||
// - If the Out type is not the empty interface [any], it provides the
|
||||
// default output schema for the tool (which again may be overridden in
|
||||
// [AddTool]).
|
||||
// - The Out value is used to populate result.StructuredOutput.
|
||||
// - If [CallToolResult.Content] is unset, it is populated with the JSON
|
||||
// content of the output.
|
||||
// - An error result is treated as a tool error, rather than a protocol
|
||||
// error, and is therefore packed into CallToolResult.Content, with
|
||||
// [IsError] set.
|
||||
//
|
||||
// For these reasons, most users can ignore the [CallToolRequest] argument and
|
||||
// [CallToolResult] return values entirely. In fact, it is permissible to
|
||||
// return a nil CallToolResult, if you only care about returning a output value
|
||||
// or error. The effective result will be populated as described above.
|
||||
type ToolHandlerFor[In, Out any] func(_ context.Context, request *CallToolRequest, input In) (result *CallToolResult, output Out, _ error)
|
||||
|
||||
// A serverTool is a tool definition that is bound to a tool handler.
|
||||
type serverTool struct {
|
||||
tool *Tool
|
||||
handler ToolHandler
|
||||
}
|
||||
|
||||
// applySchema validates whether data is valid JSON according to the provided
|
||||
// schema, after applying schema defaults.
|
||||
//
|
||||
// Returns the JSON value augmented with defaults.
|
||||
func applySchema(data json.RawMessage, resolved *jsonschema.Resolved) (json.RawMessage, error) {
|
||||
// TODO: use reflection to create the struct type to unmarshal into.
|
||||
// Separate validation from assignment.
|
||||
|
||||
// Use default JSON marshalling for validation.
|
||||
//
|
||||
// This avoids inconsistent representation due to custom marshallers, such as
|
||||
// time.Time (issue #449).
|
||||
//
|
||||
// Additionally, unmarshalling into a map ensures that the resulting JSON is
|
||||
// at least {}, even if data is empty. For example, arguments is technically
|
||||
// an optional property of callToolParams, and we still want to apply the
|
||||
// defaults in this case.
|
||||
//
|
||||
// TODO(rfindley): in which cases can resolved be nil?
|
||||
if resolved != nil {
|
||||
v := make(map[string]any)
|
||||
if len(data) > 0 {
|
||||
if err := internaljson.Unmarshal(data, &v); err != nil {
|
||||
return nil, fmt.Errorf("unmarshaling arguments: %w", err)
|
||||
}
|
||||
}
|
||||
if err := resolved.ApplyDefaults(&v); err != nil {
|
||||
return nil, fmt.Errorf("applying schema defaults:\n%w", err)
|
||||
}
|
||||
if err := resolved.Validate(&v); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// We must re-marshal with the default values applied.
|
||||
var err error
|
||||
data, err = json.Marshal(v)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshalling with defaults: %v", err)
|
||||
}
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// validateToolName checks whether name is a valid tool name, reporting a
|
||||
// non-nil error if not.
|
||||
func validateToolName(name string) error {
|
||||
if name == "" {
|
||||
return fmt.Errorf("tool name cannot be empty")
|
||||
}
|
||||
if len(name) > 128 {
|
||||
return fmt.Errorf("tool name exceeds maximum length of 128 characters (current: %d)", len(name))
|
||||
}
|
||||
// For consistency with other SDKs, report characters in the order the appear
|
||||
// in the name.
|
||||
var invalidChars []string
|
||||
seen := make(map[rune]bool)
|
||||
for _, r := range name {
|
||||
if !validToolNameRune(r) {
|
||||
if !seen[r] {
|
||||
invalidChars = append(invalidChars, fmt.Sprintf("%q", string(r)))
|
||||
seen[r] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(invalidChars) > 0 {
|
||||
return fmt.Errorf("tool name contains invalid characters: %s", strings.Join(invalidChars, ", "))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// validToolNameRune reports whether r is valid within tool names.
|
||||
func validToolNameRune(r rune) bool {
|
||||
return (r >= 'a' && r <= 'z') ||
|
||||
(r >= 'A' && r <= 'Z') ||
|
||||
(r >= '0' && r <= '9') ||
|
||||
r == '_' || r == '-' || r == '.'
|
||||
}
|
||||
+660
@@ -0,0 +1,660 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"net"
|
||||
"os"
|
||||
"sync"
|
||||
|
||||
internaljson "github.com/modelcontextprotocol/go-sdk/internal/json"
|
||||
"github.com/modelcontextprotocol/go-sdk/internal/jsonrpc2"
|
||||
"github.com/modelcontextprotocol/go-sdk/internal/xcontext"
|
||||
"github.com/modelcontextprotocol/go-sdk/jsonrpc"
|
||||
)
|
||||
|
||||
// ErrConnectionClosed is returned when sending a message to a connection that
|
||||
// is closed or in the process of closing.
|
||||
var ErrConnectionClosed = errors.New("connection closed")
|
||||
|
||||
// ErrSessionMissing is returned when the session is known to not be present on
|
||||
// the server.
|
||||
var ErrSessionMissing = errors.New("session not found")
|
||||
|
||||
// A Transport is used to create a bidirectional connection between MCP client
|
||||
// and server.
|
||||
//
|
||||
// Transports should be used for at most one call to [Server.Connect] or
|
||||
// [Client.Connect].
|
||||
type Transport interface {
|
||||
// Connect returns the logical JSON-RPC connection..
|
||||
//
|
||||
// It is called exactly once by [Server.Connect] or [Client.Connect].
|
||||
Connect(ctx context.Context) (Connection, error)
|
||||
}
|
||||
|
||||
// A Connection is a logical bidirectional JSON-RPC connection.
|
||||
type Connection interface {
|
||||
// Read reads the next message to process off the connection.
|
||||
//
|
||||
// Connections must allow Read to be called concurrently with Close. In
|
||||
// particular, calling Close should unblock a Read waiting for input.
|
||||
Read(context.Context) (jsonrpc.Message, error)
|
||||
|
||||
// Write writes a new message to the connection.
|
||||
//
|
||||
// Write may be called concurrently, as calls or responses may occur
|
||||
// concurrently in user code.
|
||||
Write(context.Context, jsonrpc.Message) error
|
||||
|
||||
// Close closes the connection. It is implicitly called whenever a Read or
|
||||
// Write fails.
|
||||
//
|
||||
// Close may be called multiple times, potentially concurrently.
|
||||
Close() error
|
||||
|
||||
// TODO(#148): remove SessionID from this interface.
|
||||
SessionID() string
|
||||
}
|
||||
|
||||
// A ClientConnection is a [Connection] that is specific to the MCP client.
|
||||
//
|
||||
// If client connections implement this interface, they may receive information
|
||||
// about changes to the client session.
|
||||
//
|
||||
// TODO: should this interface be exported?
|
||||
type clientConnection interface {
|
||||
Connection
|
||||
|
||||
// sessionUpdated is called whenever the client session state changes.
|
||||
sessionUpdated(clientSessionState)
|
||||
}
|
||||
|
||||
// A serverConnection is a Connection that is specific to the MCP server.
|
||||
//
|
||||
// If server connections implement this interface, they receive information
|
||||
// about changes to the server session.
|
||||
//
|
||||
// TODO: should this interface be exported?
|
||||
type serverConnection interface {
|
||||
Connection
|
||||
sessionUpdated(ServerSessionState)
|
||||
}
|
||||
|
||||
// A StdioTransport is a [Transport] that communicates over stdin/stdout using
|
||||
// newline-delimited JSON.
|
||||
type StdioTransport struct{}
|
||||
|
||||
// Connect implements the [Transport] interface.
|
||||
func (*StdioTransport) Connect(context.Context) (Connection, error) {
|
||||
return newIOConn(rwc{os.Stdin, nopCloserWriter{os.Stdout}}), nil
|
||||
}
|
||||
|
||||
// nopCloserWriter is an io.WriteCloser with a trivial Close method.
|
||||
type nopCloserWriter struct {
|
||||
io.Writer
|
||||
}
|
||||
|
||||
func (nopCloserWriter) Close() error { return nil }
|
||||
|
||||
// An IOTransport is a [Transport] that communicates over separate
|
||||
// io.ReadCloser and io.WriteCloser using newline-delimited JSON.
|
||||
type IOTransport struct {
|
||||
Reader io.ReadCloser
|
||||
Writer io.WriteCloser
|
||||
}
|
||||
|
||||
// Connect implements the [Transport] interface.
|
||||
func (t *IOTransport) Connect(context.Context) (Connection, error) {
|
||||
return newIOConn(rwc{t.Reader, t.Writer}), nil
|
||||
}
|
||||
|
||||
// An InMemoryTransport is a [Transport] that communicates over an in-memory
|
||||
// network connection, using newline-delimited JSON.
|
||||
//
|
||||
// InMemoryTransports should be constructed using [NewInMemoryTransports],
|
||||
// which returns two transports connected to each other.
|
||||
type InMemoryTransport struct {
|
||||
rwc io.ReadWriteCloser
|
||||
}
|
||||
|
||||
// Connect implements the [Transport] interface.
|
||||
func (t *InMemoryTransport) Connect(context.Context) (Connection, error) {
|
||||
return newIOConn(t.rwc), nil
|
||||
}
|
||||
|
||||
// NewInMemoryTransports returns two [InMemoryTransport] objects that connect
|
||||
// to each other.
|
||||
//
|
||||
// The resulting transports are symmetrical: use either to connect to a server,
|
||||
// and then the other to connect to a client. Servers must be connected before
|
||||
// clients, as the client initializes the MCP session during connection.
|
||||
func NewInMemoryTransports() (*InMemoryTransport, *InMemoryTransport) {
|
||||
c1, c2 := net.Pipe()
|
||||
return &InMemoryTransport{c1}, &InMemoryTransport{c2}
|
||||
}
|
||||
|
||||
type binder[T handler, State any] interface {
|
||||
// TODO(rfindley): the bind API has gotten too complicated. Simplify.
|
||||
bind(Connection, *jsonrpc2.Connection, State, func()) T
|
||||
disconnect(T)
|
||||
}
|
||||
|
||||
type handler interface {
|
||||
handle(ctx context.Context, req *jsonrpc.Request) (any, error)
|
||||
}
|
||||
|
||||
func connect[H handler, State any](ctx context.Context, t Transport, b binder[H, State], s State, onClose func()) (H, error) {
|
||||
var zero H
|
||||
mcpConn, err := t.Connect(ctx)
|
||||
if err != nil {
|
||||
return zero, err
|
||||
}
|
||||
// If logging is configured, write message logs.
|
||||
reader, writer := jsonrpc2.Reader(mcpConn), jsonrpc2.Writer(mcpConn)
|
||||
var (
|
||||
h H
|
||||
preempter canceller
|
||||
)
|
||||
bind := func(conn *jsonrpc2.Connection) jsonrpc2.Handler {
|
||||
h = b.bind(mcpConn, conn, s, onClose)
|
||||
preempter.conn = conn
|
||||
return jsonrpc2.HandlerFunc(h.handle)
|
||||
}
|
||||
_ = jsonrpc2.NewConnection(ctx, jsonrpc2.ConnectionConfig{
|
||||
Reader: reader,
|
||||
Writer: writer,
|
||||
Closer: mcpConn,
|
||||
Bind: bind,
|
||||
Preempter: &preempter,
|
||||
OnDone: func() {
|
||||
b.disconnect(h)
|
||||
},
|
||||
OnInternalError: func(err error) { log.Printf("jsonrpc2 error: %v", err) },
|
||||
})
|
||||
assert(preempter.conn != nil, "unbound preempter")
|
||||
return h, nil
|
||||
}
|
||||
|
||||
// A canceller is a jsonrpc2.Preempter that cancels in-flight requests on MCP
|
||||
// cancelled notifications.
|
||||
type canceller struct {
|
||||
conn *jsonrpc2.Connection
|
||||
}
|
||||
|
||||
// Preempt implements [jsonrpc2.Preempter].
|
||||
func (c *canceller) Preempt(ctx context.Context, req *jsonrpc.Request) (result any, err error) {
|
||||
if req.Method == notificationCancelled {
|
||||
var params CancelledParams
|
||||
if err := internaljson.Unmarshal(req.Params, ¶ms); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
id, err := jsonrpc2.MakeID(params.RequestID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
go c.conn.Cancel(id)
|
||||
}
|
||||
return nil, jsonrpc2.ErrNotHandled
|
||||
}
|
||||
|
||||
// call executes and awaits a jsonrpc2 call on the given connection,
|
||||
// translating errors into the mcp domain.
|
||||
func call(ctx context.Context, conn *jsonrpc2.Connection, method string, params Params, result Result) error {
|
||||
// The "%w"s in this function expose jsonrpc.Error as part of the API.
|
||||
call := conn.Call(ctx, method, params)
|
||||
err := call.Await(ctx, result)
|
||||
switch {
|
||||
case errors.Is(err, jsonrpc2.ErrClientClosing), errors.Is(err, jsonrpc2.ErrServerClosing):
|
||||
return fmt.Errorf("%w: calling %q: %v", ErrConnectionClosed, method, err)
|
||||
case ctx.Err() != nil:
|
||||
// Notify the peer of cancellation.
|
||||
err := conn.Notify(xcontext.Detach(ctx), notificationCancelled, &CancelledParams{
|
||||
Reason: ctx.Err().Error(),
|
||||
RequestID: call.ID().Raw(),
|
||||
})
|
||||
// By default, the jsonrpc2 library waits for graceful shutdown when the
|
||||
// connection is closed, meaning it expects all outgoing and incoming
|
||||
// requests to complete. However, for MCP this expectation is unrealistic,
|
||||
// and can lead to hanging shutdown. For example, if a streamable client is
|
||||
// killed, the server will not be able to detect this event, except via
|
||||
// keepalive pings (if they are configured), and so outgoing calls may hang
|
||||
// indefinitely.
|
||||
//
|
||||
// Therefore, we choose to eagerly retire calls, removing them from the
|
||||
// outgoingCalls map, when the caller context is cancelled: if the caller
|
||||
// will never receive the response, there's no need to track it.
|
||||
conn.Retire(call, ctx.Err())
|
||||
return errors.Join(ctx.Err(), err)
|
||||
case err != nil:
|
||||
return fmt.Errorf("calling %q: %w", method, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// A LoggingTransport is a [Transport] that delegates to another transport,
|
||||
// writing RPC logs to an io.Writer.
|
||||
type LoggingTransport struct {
|
||||
Transport Transport
|
||||
Writer io.Writer
|
||||
}
|
||||
|
||||
// Connect connects the underlying transport, returning a [Connection] that writes
|
||||
// logs to the configured destination.
|
||||
func (t *LoggingTransport) Connect(ctx context.Context) (Connection, error) {
|
||||
delegate, err := t.Transport.Connect(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &loggingConn{delegate: delegate, w: t.Writer}, nil
|
||||
}
|
||||
|
||||
type loggingConn struct {
|
||||
delegate Connection
|
||||
|
||||
mu sync.Mutex
|
||||
w io.Writer
|
||||
}
|
||||
|
||||
func (c *loggingConn) SessionID() string { return c.delegate.SessionID() }
|
||||
|
||||
// Read is a stream middleware that logs incoming messages.
|
||||
func (s *loggingConn) Read(ctx context.Context) (jsonrpc.Message, error) {
|
||||
msg, err := s.delegate.Read(ctx)
|
||||
|
||||
if err != nil {
|
||||
s.mu.Lock()
|
||||
fmt.Fprintf(s.w, "read error: %v\n", err)
|
||||
s.mu.Unlock()
|
||||
} else {
|
||||
data, err := jsonrpc2.EncodeMessage(msg)
|
||||
s.mu.Lock()
|
||||
if err != nil {
|
||||
fmt.Fprintf(s.w, "LoggingTransport: failed to marshal: %v", err)
|
||||
}
|
||||
fmt.Fprintf(s.w, "read: %s\n", string(data))
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
return msg, err
|
||||
}
|
||||
|
||||
// Write is a stream middleware that logs outgoing messages.
|
||||
func (s *loggingConn) Write(ctx context.Context, msg jsonrpc.Message) error {
|
||||
err := s.delegate.Write(ctx, msg)
|
||||
if err != nil {
|
||||
s.mu.Lock()
|
||||
fmt.Fprintf(s.w, "write error: %v\n", err)
|
||||
s.mu.Unlock()
|
||||
} else {
|
||||
data, err := jsonrpc2.EncodeMessage(msg)
|
||||
s.mu.Lock()
|
||||
if err != nil {
|
||||
fmt.Fprintf(s.w, "LoggingTransport: failed to marshal: %v", err)
|
||||
}
|
||||
fmt.Fprintf(s.w, "write: %s\n", string(data))
|
||||
s.mu.Unlock()
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *loggingConn) Close() error {
|
||||
return s.delegate.Close()
|
||||
}
|
||||
|
||||
// A rwc binds an io.ReadCloser and io.WriteCloser together to create an
|
||||
// io.ReadWriteCloser.
|
||||
type rwc struct {
|
||||
rc io.ReadCloser
|
||||
wc io.WriteCloser
|
||||
}
|
||||
|
||||
func (r rwc) Read(p []byte) (n int, err error) {
|
||||
return r.rc.Read(p)
|
||||
}
|
||||
|
||||
func (r rwc) Write(p []byte) (n int, err error) {
|
||||
return r.wc.Write(p)
|
||||
}
|
||||
|
||||
func (r rwc) Close() error {
|
||||
rcErr := r.rc.Close()
|
||||
|
||||
var wcErr error
|
||||
if r.wc != nil { // we only allow a nil writer in unit tests
|
||||
wcErr = r.wc.Close()
|
||||
}
|
||||
|
||||
return errors.Join(rcErr, wcErr)
|
||||
}
|
||||
|
||||
// An ioConn is a transport that delimits messages with newlines across
|
||||
// a bidirectional stream, and supports jsonrpc.2 message batching.
|
||||
//
|
||||
// See https://github.com/ndjson/ndjson-spec for discussion of newline
|
||||
// delimited JSON.
|
||||
//
|
||||
// See [msgBatch] for more discussion of message batching.
|
||||
type ioConn struct {
|
||||
protocolVersion string // negotiated version, set during session initialization.
|
||||
|
||||
writeMu sync.Mutex // guards Write, which must be concurrency safe.
|
||||
rwc io.ReadWriteCloser // the underlying stream
|
||||
|
||||
// incoming receives messages from the read loop started in [newIOConn].
|
||||
incoming <-chan msgOrErr
|
||||
|
||||
// If outgoiBatch has a positive capacity, it will be used to batch requests
|
||||
// and notifications before sending.
|
||||
outgoingBatch []jsonrpc.Message
|
||||
|
||||
// Unread messages in the last batch. Since reads are serialized, there is no
|
||||
// need to guard here.
|
||||
queue []jsonrpc.Message
|
||||
|
||||
// batches correlate incoming requests to the batch in which they arrived.
|
||||
// Since writes may be concurrent to reads, we need to guard this with a mutex.
|
||||
batchMu sync.Mutex
|
||||
batches map[jsonrpc2.ID]*msgBatch // lazily allocated
|
||||
|
||||
closeOnce sync.Once
|
||||
closed chan struct{}
|
||||
closeErr error
|
||||
}
|
||||
|
||||
type msgOrErr struct {
|
||||
msg json.RawMessage
|
||||
err error
|
||||
}
|
||||
|
||||
func newIOConn(rwc io.ReadWriteCloser) *ioConn {
|
||||
var (
|
||||
incoming = make(chan msgOrErr)
|
||||
closed = make(chan struct{})
|
||||
)
|
||||
// Start a goroutine for reads, so that we can select on the incoming channel
|
||||
// in [ioConn.Read] and unblock the read as soon as Close is called (see #224).
|
||||
//
|
||||
// This leaks a goroutine if rwc.Read does not unblock after it is closed,
|
||||
// but that is unavoidable since AFAIK there is no (easy and portable) way to
|
||||
// guarantee that reads of stdin are unblocked when closed.
|
||||
go func() {
|
||||
dec := json.NewDecoder(rwc)
|
||||
for {
|
||||
var raw json.RawMessage
|
||||
err := dec.Decode(&raw)
|
||||
// If decoding was successful, check for trailing data at the end of the stream.
|
||||
if err == nil {
|
||||
// Read the next byte to check if there is trailing data.
|
||||
var tr [1]byte
|
||||
if n, readErr := dec.Buffered().Read(tr[:]); n > 0 {
|
||||
// If read byte is not a newline, it is an error.
|
||||
// Support both Unix (\n) and Windows (\r\n) line endings.
|
||||
if tr[0] != '\n' && tr[0] != '\r' {
|
||||
err = fmt.Errorf("invalid trailing data at the end of stream")
|
||||
}
|
||||
} else if readErr != nil && readErr != io.EOF {
|
||||
err = readErr
|
||||
}
|
||||
}
|
||||
select {
|
||||
case incoming <- msgOrErr{msg: raw, err: err}:
|
||||
case <-closed:
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
return &ioConn{
|
||||
rwc: rwc,
|
||||
incoming: incoming,
|
||||
closed: closed,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *ioConn) SessionID() string { return "" }
|
||||
|
||||
func (c *ioConn) sessionUpdated(state ServerSessionState) {
|
||||
protocolVersion := ""
|
||||
if state.InitializeParams != nil {
|
||||
protocolVersion = state.InitializeParams.ProtocolVersion
|
||||
}
|
||||
if protocolVersion == "" {
|
||||
protocolVersion = protocolVersion20250326
|
||||
}
|
||||
c.protocolVersion = negotiatedVersion(protocolVersion)
|
||||
}
|
||||
|
||||
// addBatch records a msgBatch for an incoming batch payload.
|
||||
// It returns an error if batch is malformed, containing previously seen IDs.
|
||||
//
|
||||
// See [msgBatch] for more.
|
||||
func (t *ioConn) addBatch(batch *msgBatch) error {
|
||||
t.batchMu.Lock()
|
||||
defer t.batchMu.Unlock()
|
||||
for id := range batch.unresolved {
|
||||
if _, ok := t.batches[id]; ok {
|
||||
return fmt.Errorf("%w: batch contains previously seen request %v", jsonrpc2.ErrInvalidRequest, id.Raw())
|
||||
}
|
||||
}
|
||||
for id := range batch.unresolved {
|
||||
if t.batches == nil {
|
||||
t.batches = make(map[jsonrpc2.ID]*msgBatch)
|
||||
}
|
||||
t.batches[id] = batch
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// updateBatch records a response in the message batch tracking the
|
||||
// corresponding incoming call, if any.
|
||||
//
|
||||
// The second result reports whether resp was part of a batch. If this is true,
|
||||
// the first result is nil if the batch is still incomplete, or the full set of
|
||||
// batch responses if resp completed the batch.
|
||||
func (t *ioConn) updateBatch(resp *jsonrpc.Response) ([]*jsonrpc.Response, bool) {
|
||||
t.batchMu.Lock()
|
||||
defer t.batchMu.Unlock()
|
||||
|
||||
if batch, ok := t.batches[resp.ID]; ok {
|
||||
idx, ok := batch.unresolved[resp.ID]
|
||||
if !ok {
|
||||
panic("internal error: inconsistent batches")
|
||||
}
|
||||
batch.responses[idx] = resp
|
||||
delete(batch.unresolved, resp.ID)
|
||||
delete(t.batches, resp.ID)
|
||||
if len(batch.unresolved) == 0 {
|
||||
return batch.responses, true
|
||||
}
|
||||
return nil, true
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// A msgBatch records information about an incoming batch of jsonrpc.2 calls.
|
||||
//
|
||||
// The jsonrpc.2 spec (https://www.jsonrpc.org/specification#batch) says:
|
||||
//
|
||||
// "The Server should respond with an Array containing the corresponding
|
||||
// Response objects, after all of the batch Request objects have been
|
||||
// processed. A Response object SHOULD exist for each Request object, except
|
||||
// that there SHOULD NOT be any Response objects for notifications. The Server
|
||||
// MAY process a batch rpc call as a set of concurrent tasks, processing them
|
||||
// in any order and with any width of parallelism."
|
||||
//
|
||||
// Therefore, a msgBatch keeps track of outstanding calls and their responses.
|
||||
// When there are no unresolved calls, the response payload is sent.
|
||||
type msgBatch struct {
|
||||
unresolved map[jsonrpc2.ID]int
|
||||
responses []*jsonrpc.Response
|
||||
}
|
||||
|
||||
func (t *ioConn) Read(ctx context.Context) (jsonrpc.Message, error) {
|
||||
// As a matter of principle, enforce that reads on a closed context return an
|
||||
// error.
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
default:
|
||||
}
|
||||
if len(t.queue) > 0 {
|
||||
next := t.queue[0]
|
||||
t.queue = t.queue[1:]
|
||||
return next, nil
|
||||
}
|
||||
|
||||
var raw json.RawMessage
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
|
||||
case v := <-t.incoming:
|
||||
if v.err != nil {
|
||||
return nil, v.err
|
||||
}
|
||||
raw = v.msg
|
||||
|
||||
case <-t.closed:
|
||||
return nil, io.EOF
|
||||
}
|
||||
|
||||
msgs, batch, err := readBatch(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if batch && t.protocolVersion >= protocolVersion20250618 {
|
||||
return nil, fmt.Errorf("JSON-RPC batching is not supported in %s and later (request version: %s)", protocolVersion20250618, t.protocolVersion)
|
||||
}
|
||||
|
||||
t.queue = msgs[1:]
|
||||
|
||||
if batch {
|
||||
var respBatch *msgBatch // track incoming requests in the batch
|
||||
for _, msg := range msgs {
|
||||
if req, ok := msg.(*jsonrpc.Request); ok {
|
||||
if respBatch == nil {
|
||||
respBatch = &msgBatch{
|
||||
unresolved: make(map[jsonrpc2.ID]int),
|
||||
}
|
||||
}
|
||||
if _, ok := respBatch.unresolved[req.ID]; ok {
|
||||
return nil, fmt.Errorf("duplicate message ID %q", req.ID)
|
||||
}
|
||||
respBatch.unresolved[req.ID] = len(respBatch.responses)
|
||||
respBatch.responses = append(respBatch.responses, nil)
|
||||
}
|
||||
}
|
||||
if respBatch != nil {
|
||||
// The batch contains one or more incoming requests to track.
|
||||
if err := t.addBatch(respBatch); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
}
|
||||
return msgs[0], err
|
||||
}
|
||||
|
||||
// readBatch reads batch data, which may be either a single JSON-RPC message,
|
||||
// or an array of JSON-RPC messages.
|
||||
func readBatch(data []byte) (msgs []jsonrpc.Message, isBatch bool, _ error) {
|
||||
// Try to read an array of messages first.
|
||||
var rawBatch []json.RawMessage
|
||||
if err := internaljson.Unmarshal(data, &rawBatch); err == nil {
|
||||
if len(rawBatch) == 0 {
|
||||
return nil, true, fmt.Errorf("empty batch")
|
||||
}
|
||||
for _, raw := range rawBatch {
|
||||
msg, err := jsonrpc2.DecodeMessage(raw)
|
||||
if err != nil {
|
||||
return nil, true, err
|
||||
}
|
||||
msgs = append(msgs, msg)
|
||||
}
|
||||
return msgs, true, nil
|
||||
}
|
||||
// Try again with a single message.
|
||||
msg, err := jsonrpc2.DecodeMessage(data)
|
||||
return []jsonrpc.Message{msg}, false, err
|
||||
}
|
||||
|
||||
func (t *ioConn) Write(ctx context.Context, msg jsonrpc.Message) error {
|
||||
// As in [ioConn.Read], enforce that Writes on a closed context are an error.
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
default:
|
||||
}
|
||||
|
||||
t.writeMu.Lock()
|
||||
defer t.writeMu.Unlock()
|
||||
|
||||
// Batching support: if msg is a Response, it may have completed a batch, so
|
||||
// check that first. Otherwise, it is a request or notification, and we may
|
||||
// want to collect it into a batch before sending, if we're configured to use
|
||||
// outgoing batches.
|
||||
if resp, ok := msg.(*jsonrpc.Response); ok {
|
||||
if batch, ok := t.updateBatch(resp); ok {
|
||||
if len(batch) > 0 {
|
||||
data, err := marshalMessages(batch)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
data = append(data, '\n')
|
||||
_, err = t.rwc.Write(data)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
} else if len(t.outgoingBatch) < cap(t.outgoingBatch) {
|
||||
t.outgoingBatch = append(t.outgoingBatch, msg)
|
||||
if len(t.outgoingBatch) == cap(t.outgoingBatch) {
|
||||
data, err := marshalMessages(t.outgoingBatch)
|
||||
t.outgoingBatch = t.outgoingBatch[:0]
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
data = append(data, '\n')
|
||||
_, err = t.rwc.Write(data)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
data, err := jsonrpc2.EncodeMessage(msg)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshaling message: %v", err)
|
||||
}
|
||||
data = append(data, '\n') // newline delimited
|
||||
_, err = t.rwc.Write(data)
|
||||
return err
|
||||
}
|
||||
|
||||
func (t *ioConn) Close() error {
|
||||
t.closeOnce.Do(func() {
|
||||
t.closeErr = t.rwc.Close()
|
||||
close(t.closed)
|
||||
})
|
||||
return t.closeErr
|
||||
}
|
||||
|
||||
func marshalMessages[T jsonrpc.Message](msgs []T) ([]byte, error) {
|
||||
var rawMsgs []json.RawMessage
|
||||
for _, msg := range msgs {
|
||||
raw, err := jsonrpc2.EncodeMessage(msg)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("encoding batch message: %w", err)
|
||||
}
|
||||
rawMsgs = append(rawMsgs, raw)
|
||||
}
|
||||
return json.Marshal(rawMsgs)
|
||||
}
|
||||
+30
@@ -0,0 +1,30 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
|
||||
internaljson "github.com/modelcontextprotocol/go-sdk/internal/json"
|
||||
)
|
||||
|
||||
func assert(cond bool, msg string) {
|
||||
if !cond {
|
||||
panic(msg)
|
||||
}
|
||||
}
|
||||
|
||||
// remarshal marshals from to JSON, and then unmarshals into to, which must be
|
||||
// a pointer type.
|
||||
func remarshal(from, to any) error {
|
||||
data, err := json.Marshal(from)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := internaljson.Unmarshal(data, to); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
+210
@@ -0,0 +1,210 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// This file implements Authorization Server Metadata.
|
||||
// See https://www.rfc-editor.org/rfc/rfc8414.html.
|
||||
|
||||
//go:build mcp_go_client_oauth
|
||||
|
||||
package oauthex
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
// AuthServerMeta represents the metadata for an OAuth 2.0 authorization server,
|
||||
// as defined in [RFC 8414].
|
||||
//
|
||||
// Not supported:
|
||||
// - signed metadata
|
||||
//
|
||||
// Note: URL fields in this struct are validated by validateAuthServerMetaURLs to
|
||||
// prevent XSS attacks. If you add a new URL field, you must also add it to that
|
||||
// function.
|
||||
//
|
||||
// [RFC 8414]: https://tools.ietf.org/html/rfc8414)
|
||||
type AuthServerMeta struct {
|
||||
// Issuer is the REQUIRED URL identifying the authorization server.
|
||||
Issuer string `json:"issuer"`
|
||||
|
||||
// AuthorizationEndpoint is the REQUIRED URL of the server's OAuth 2.0 authorization endpoint.
|
||||
AuthorizationEndpoint string `json:"authorization_endpoint"`
|
||||
|
||||
// TokenEndpoint is the REQUIRED URL of the server's OAuth 2.0 token endpoint.
|
||||
TokenEndpoint string `json:"token_endpoint"`
|
||||
|
||||
// JWKSURI is the REQUIRED URL of the server's JSON Web Key Set [JWK] document.
|
||||
JWKSURI string `json:"jwks_uri"`
|
||||
|
||||
// RegistrationEndpoint is the RECOMMENDED URL of the server's OAuth 2.0 Dynamic Client Registration endpoint.
|
||||
RegistrationEndpoint string `json:"registration_endpoint,omitempty"`
|
||||
|
||||
// ScopesSupported is a RECOMMENDED JSON array of strings containing a list of the OAuth 2.0
|
||||
// "scope" values that this server supports.
|
||||
ScopesSupported []string `json:"scopes_supported,omitempty"`
|
||||
|
||||
// ResponseTypesSupported is a REQUIRED JSON array of strings containing a list of the OAuth 2.0
|
||||
// "response_type" values that this server supports.
|
||||
ResponseTypesSupported []string `json:"response_types_supported"`
|
||||
|
||||
// ResponseModesSupported is a RECOMMENDED JSON array of strings containing a list of the OAuth 2.0
|
||||
// "response_mode" values that this server supports.
|
||||
ResponseModesSupported []string `json:"response_modes_supported,omitempty"`
|
||||
|
||||
// GrantTypesSupported is a RECOMMENDED JSON array of strings containing a list of the OAuth 2.0
|
||||
// grant type values that this server supports.
|
||||
GrantTypesSupported []string `json:"grant_types_supported,omitempty"`
|
||||
|
||||
// TokenEndpointAuthMethodsSupported is a RECOMMENDED JSON array of strings containing a list of
|
||||
// client authentication methods supported by this token endpoint.
|
||||
TokenEndpointAuthMethodsSupported []string `json:"token_endpoint_auth_methods_supported,omitempty"`
|
||||
|
||||
// TokenEndpointAuthSigningAlgValuesSupported is a RECOMMENDED JSON array of strings containing
|
||||
// a list of the JWS signing algorithms ("alg" values) supported by the token endpoint for
|
||||
// the signature on the JWT used to authenticate the client.
|
||||
TokenEndpointAuthSigningAlgValuesSupported []string `json:"token_endpoint_auth_signing_alg_values_supported,omitempty"`
|
||||
|
||||
// ServiceDocumentation is a RECOMMENDED URL of a page containing human-readable documentation
|
||||
// for the service.
|
||||
ServiceDocumentation string `json:"service_documentation,omitempty"`
|
||||
|
||||
// UILocalesSupported is a RECOMMENDED JSON array of strings representing supported
|
||||
// BCP47 [RFC5646] language tag values for display in the user interface.
|
||||
UILocalesSupported []string `json:"ui_locales_supported,omitempty"`
|
||||
|
||||
// OpPolicyURI is a RECOMMENDED URL that the server provides to the person registering
|
||||
// the client to read about the server's operator policies.
|
||||
OpPolicyURI string `json:"op_policy_uri,omitempty"`
|
||||
|
||||
// OpTOSURI is a RECOMMENDED URL that the server provides to the person registering the
|
||||
// client to read about the server's terms of service.
|
||||
OpTOSURI string `json:"op_tos_uri,omitempty"`
|
||||
|
||||
// RevocationEndpoint is a RECOMMENDED URL of the server's OAuth 2.0 revocation endpoint.
|
||||
RevocationEndpoint string `json:"revocation_endpoint,omitempty"`
|
||||
|
||||
// RevocationEndpointAuthMethodsSupported is a RECOMMENDED JSON array of strings containing
|
||||
// a list of client authentication methods supported by this revocation endpoint.
|
||||
RevocationEndpointAuthMethodsSupported []string `json:"revocation_endpoint_auth_methods_supported,omitempty"`
|
||||
|
||||
// RevocationEndpointAuthSigningAlgValuesSupported is a RECOMMENDED JSON array of strings
|
||||
// containing a list of the JWS signing algorithms ("alg" values) supported by the revocation
|
||||
// endpoint for the signature on the JWT used to authenticate the client.
|
||||
RevocationEndpointAuthSigningAlgValuesSupported []string `json:"revocation_endpoint_auth_signing_alg_values_supported,omitempty"`
|
||||
|
||||
// IntrospectionEndpoint is a RECOMMENDED URL of the server's OAuth 2.0 introspection endpoint.
|
||||
IntrospectionEndpoint string `json:"introspection_endpoint,omitempty"`
|
||||
|
||||
// IntrospectionEndpointAuthMethodsSupported is a RECOMMENDED JSON array of strings containing
|
||||
// a list of client authentication methods supported by this introspection endpoint.
|
||||
IntrospectionEndpointAuthMethodsSupported []string `json:"introspection_endpoint_auth_methods_supported,omitempty"`
|
||||
|
||||
// IntrospectionEndpointAuthSigningAlgValuesSupported is a RECOMMENDED JSON array of strings
|
||||
// containing a list of the JWS signing algorithms ("alg" values) supported by the introspection
|
||||
// endpoint for the signature on the JWT used to authenticate the client.
|
||||
IntrospectionEndpointAuthSigningAlgValuesSupported []string `json:"introspection_endpoint_auth_signing_alg_values_supported,omitempty"`
|
||||
|
||||
// CodeChallengeMethodsSupported is a RECOMMENDED JSON array of strings containing a list of
|
||||
// PKCE code challenge methods supported by this authorization server.
|
||||
CodeChallengeMethodsSupported []string `json:"code_challenge_methods_supported,omitempty"`
|
||||
|
||||
// ClientIDMetadataDocumentSupported is a boolean indicating whether the authorization server
|
||||
// supports client ID metadata documents.
|
||||
ClientIDMetadataDocumentSupported bool `json:"client_id_metadata_document_supported,omitempty"`
|
||||
}
|
||||
|
||||
// GetAuthServerMeta issues a GET request to retrieve authorization server metadata
|
||||
// from an OAuth authorization server with the given metadataURL.
|
||||
//
|
||||
// It follows [RFC 8414]:
|
||||
// - The metadataURL must use HTTPS or be a local address.
|
||||
// - The Issuer field is checked against metadataURL.Issuer.
|
||||
//
|
||||
// It also verifies that the authorization server supports PKCE and that the URLs
|
||||
// in the metadata don't use dangerous schemes.
|
||||
//
|
||||
// It returns an error if the request fails with a non-4xx status code or the fetched
|
||||
// metadata doesn't pass security validations.
|
||||
// It returns nil if the request fails with a 4xx status code.
|
||||
//
|
||||
// [RFC 8414]: https://tools.ietf.org/html/rfc8414
|
||||
func GetAuthServerMeta(ctx context.Context, metadataURL, issuer string, c *http.Client) (*AuthServerMeta, error) {
|
||||
// Only allow HTTP for local addresses (testing or development purposes).
|
||||
if err := checkHTTPSOrLoopback(metadataURL); err != nil {
|
||||
return nil, fmt.Errorf("metadataURL: %v", err)
|
||||
}
|
||||
asm, err := getJSON[AuthServerMeta](ctx, c, metadataURL, 1<<20)
|
||||
if err != nil {
|
||||
var httpErr *httpStatusError
|
||||
if errors.As(err, &httpErr) {
|
||||
if 400 <= httpErr.StatusCode && httpErr.StatusCode < 500 {
|
||||
return nil, nil
|
||||
}
|
||||
}
|
||||
return nil, fmt.Errorf("%v", err) // Do not expose error types.
|
||||
}
|
||||
if asm.Issuer != issuer {
|
||||
// Validate the Issuer field (see RFC 8414, section 3.3).
|
||||
return nil, fmt.Errorf("metadata issuer %q does not match issuer URL %q", asm.Issuer, issuer)
|
||||
}
|
||||
|
||||
if len(asm.CodeChallengeMethodsSupported) == 0 {
|
||||
return nil, fmt.Errorf("authorization server at %s does not implement PKCE", issuer)
|
||||
}
|
||||
|
||||
// Validate endpoint URLs to prevent XSS attacks (see #526).
|
||||
if err := validateAuthServerMetaURLs(asm); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return asm, nil
|
||||
}
|
||||
|
||||
// validateAuthServerMetaURLs validates all URL fields in AuthServerMeta
|
||||
// to ensure they don't use dangerous schemes that could enable XSS attacks.
|
||||
// It also validates that URLs likely to be called by the client use
|
||||
// HTTPS or are loopback addresses.
|
||||
func validateAuthServerMetaURLs(asm *AuthServerMeta) error {
|
||||
urls := []struct {
|
||||
name string
|
||||
value string
|
||||
}{
|
||||
{"authorization_endpoint", asm.AuthorizationEndpoint},
|
||||
{"token_endpoint", asm.TokenEndpoint},
|
||||
{"jwks_uri", asm.JWKSURI},
|
||||
{"registration_endpoint", asm.RegistrationEndpoint},
|
||||
{"service_documentation", asm.ServiceDocumentation},
|
||||
{"op_policy_uri", asm.OpPolicyURI},
|
||||
{"op_tos_uri", asm.OpTOSURI},
|
||||
{"revocation_endpoint", asm.RevocationEndpoint},
|
||||
{"introspection_endpoint", asm.IntrospectionEndpoint},
|
||||
}
|
||||
|
||||
for _, u := range urls {
|
||||
if err := checkURLScheme(u.value); err != nil {
|
||||
return fmt.Errorf("%s: %w", u.name, err)
|
||||
}
|
||||
}
|
||||
|
||||
urls = []struct {
|
||||
name string
|
||||
value string
|
||||
}{
|
||||
{"authorization_endpoint", asm.AuthorizationEndpoint},
|
||||
{"token_endpoint", asm.TokenEndpoint},
|
||||
{"registration_endpoint", asm.RegistrationEndpoint},
|
||||
{"introspection_endpoint", asm.IntrospectionEndpoint},
|
||||
}
|
||||
|
||||
for _, u := range urls {
|
||||
if err := checkHTTPSOrLoopback(u.value); err != nil {
|
||||
return fmt.Errorf("%s: %w", u.name, err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
+263
@@ -0,0 +1,263 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// This file implements Authorization Server Metadata.
|
||||
// See https://www.rfc-editor.org/rfc/rfc8414.html.
|
||||
|
||||
//go:build mcp_go_client_oauth
|
||||
|
||||
package oauthex
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
internaljson "github.com/modelcontextprotocol/go-sdk/internal/json"
|
||||
)
|
||||
|
||||
// ClientRegistrationMetadata represents the client metadata fields for the DCR POST request (RFC 7591).
|
||||
//
|
||||
// Note: URL fields in this struct are validated by validateClientRegistrationURLs
|
||||
// to prevent XSS attacks. If you add a new URL field, you must also add it to
|
||||
// that function.
|
||||
type ClientRegistrationMetadata struct {
|
||||
// RedirectURIs is a REQUIRED JSON array of redirection URI strings for use in
|
||||
// redirect-based flows (such as the authorization code grant).
|
||||
RedirectURIs []string `json:"redirect_uris"`
|
||||
|
||||
// TokenEndpointAuthMethod is an OPTIONAL string indicator of the requested
|
||||
// authentication method for the token endpoint.
|
||||
// If omitted, the default is "client_secret_basic".
|
||||
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"`
|
||||
|
||||
// GrantTypes is an OPTIONAL JSON array of OAuth 2.0 grant type strings
|
||||
// that the client will restrict itself to using.
|
||||
// If omitted, the default is ["authorization_code"].
|
||||
GrantTypes []string `json:"grant_types,omitempty"`
|
||||
|
||||
// ResponseTypes is an OPTIONAL JSON array of OAuth 2.0 response type strings
|
||||
// that the client will restrict itself to using.
|
||||
// If omitted, the default is ["code"].
|
||||
ResponseTypes []string `json:"response_types,omitempty"`
|
||||
|
||||
// ClientName is a RECOMMENDED human-readable name of the client to be presented
|
||||
// to the end-user.
|
||||
ClientName string `json:"client_name,omitempty"`
|
||||
|
||||
// ClientURI is a RECOMMENDED URL of a web page providing information about the client.
|
||||
ClientURI string `json:"client_uri,omitempty"`
|
||||
|
||||
// LogoURI is an OPTIONAL URL of a logo for the client, which may be displayed
|
||||
// to the end-user.
|
||||
LogoURI string `json:"logo_uri,omitempty"`
|
||||
|
||||
// Scope is an OPTIONAL string containing a space-separated list of scope values
|
||||
// that the client will restrict itself to using.
|
||||
Scope string `json:"scope,omitempty"`
|
||||
|
||||
// Contacts is an OPTIONAL JSON array of strings representing ways to contact
|
||||
// people responsible for this client (e.g., email addresses).
|
||||
Contacts []string `json:"contacts,omitempty"`
|
||||
|
||||
// TOSURI is an OPTIONAL URL that the client provides to the end-user
|
||||
// to read about the client's terms of service.
|
||||
TOSURI string `json:"tos_uri,omitempty"`
|
||||
|
||||
// PolicyURI is an OPTIONAL URL that the client provides to the end-user
|
||||
// to read about the client's privacy policy.
|
||||
PolicyURI string `json:"policy_uri,omitempty"`
|
||||
|
||||
// JWKSURI is an OPTIONAL URL for the client's JSON Web Key Set [JWK] document.
|
||||
// This is preferred over the 'jwks' parameter.
|
||||
JWKSURI string `json:"jwks_uri,omitempty"`
|
||||
|
||||
// JWKS is an OPTIONAL client's JSON Web Key Set [JWK] document, passed by value.
|
||||
// This is an alternative to providing a JWKSURI.
|
||||
JWKS string `json:"jwks,omitempty"`
|
||||
|
||||
// SoftwareID is an OPTIONAL unique identifier string for the client software,
|
||||
// constant across all instances and versions.
|
||||
SoftwareID string `json:"software_id,omitempty"`
|
||||
|
||||
// SoftwareVersion is an OPTIONAL version identifier string for the client software.
|
||||
SoftwareVersion string `json:"software_version,omitempty"`
|
||||
|
||||
// SoftwareStatement is an OPTIONAL JWT that asserts client metadata values.
|
||||
// Values in the software statement take precedence over other metadata values.
|
||||
SoftwareStatement string `json:"software_statement,omitempty"`
|
||||
}
|
||||
|
||||
// ClientRegistrationResponse represents the fields returned by the Authorization Server
|
||||
// (RFC 7591, Section 3.2.1 and 3.2.2).
|
||||
type ClientRegistrationResponse struct {
|
||||
// ClientRegistrationMetadata contains all registered client metadata, returned by the
|
||||
// server on success, potentially with modified or defaulted values.
|
||||
ClientRegistrationMetadata
|
||||
|
||||
// ClientID is the REQUIRED newly issued OAuth 2.0 client identifier.
|
||||
ClientID string `json:"client_id"`
|
||||
|
||||
// ClientSecret is an OPTIONAL client secret string.
|
||||
ClientSecret string `json:"client_secret,omitempty"`
|
||||
|
||||
// ClientIDIssuedAt is an OPTIONAL Unix timestamp when the ClientID was issued.
|
||||
ClientIDIssuedAt time.Time `json:"client_id_issued_at,omitempty"`
|
||||
|
||||
// ClientSecretExpiresAt is the REQUIRED (if client_secret is issued) Unix
|
||||
// timestamp when the secret expires, or 0 if it never expires.
|
||||
ClientSecretExpiresAt time.Time `json:"client_secret_expires_at,omitempty"`
|
||||
}
|
||||
|
||||
func (r *ClientRegistrationResponse) MarshalJSON() ([]byte, error) {
|
||||
type alias ClientRegistrationResponse
|
||||
var clientIDIssuedAt int64
|
||||
var clientSecretExpiresAt int64
|
||||
|
||||
if !r.ClientIDIssuedAt.IsZero() {
|
||||
clientIDIssuedAt = r.ClientIDIssuedAt.Unix()
|
||||
}
|
||||
if !r.ClientSecretExpiresAt.IsZero() {
|
||||
clientSecretExpiresAt = r.ClientSecretExpiresAt.Unix()
|
||||
}
|
||||
|
||||
return json.Marshal(&struct {
|
||||
ClientIDIssuedAt int64 `json:"client_id_issued_at,omitempty"`
|
||||
ClientSecretExpiresAt int64 `json:"client_secret_expires_at,omitempty"`
|
||||
*alias
|
||||
}{
|
||||
ClientIDIssuedAt: clientIDIssuedAt,
|
||||
ClientSecretExpiresAt: clientSecretExpiresAt,
|
||||
alias: (*alias)(r),
|
||||
})
|
||||
}
|
||||
|
||||
func (r *ClientRegistrationResponse) UnmarshalJSON(data []byte) error {
|
||||
type alias ClientRegistrationResponse
|
||||
aux := &struct {
|
||||
ClientIDIssuedAt int64 `json:"client_id_issued_at,omitempty"`
|
||||
ClientSecretExpiresAt int64 `json:"client_secret_expires_at,omitempty"`
|
||||
*alias
|
||||
}{
|
||||
alias: (*alias)(r),
|
||||
}
|
||||
if err := internaljson.Unmarshal(data, &aux); err != nil {
|
||||
return err
|
||||
}
|
||||
if aux.ClientIDIssuedAt != 0 {
|
||||
r.ClientIDIssuedAt = time.Unix(aux.ClientIDIssuedAt, 0)
|
||||
}
|
||||
if aux.ClientSecretExpiresAt != 0 {
|
||||
r.ClientSecretExpiresAt = time.Unix(aux.ClientSecretExpiresAt, 0)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ClientRegistrationError is the error response from the Authorization Server
|
||||
// for a failed registration attempt (RFC 7591, Section 3.2.2).
|
||||
type ClientRegistrationError struct {
|
||||
// ErrorCode is the REQUIRED error code if registration failed (RFC 7591, 3.2.2).
|
||||
ErrorCode string `json:"error"`
|
||||
|
||||
// ErrorDescription is an OPTIONAL human-readable error message.
|
||||
ErrorDescription string `json:"error_description,omitempty"`
|
||||
}
|
||||
|
||||
func (e *ClientRegistrationError) Error() string {
|
||||
return fmt.Sprintf("registration failed: %s (%s)", e.ErrorCode, e.ErrorDescription)
|
||||
}
|
||||
|
||||
// RegisterClient performs Dynamic Client Registration according to RFC 7591.
|
||||
func RegisterClient(ctx context.Context, registrationEndpoint string, clientMeta *ClientRegistrationMetadata, c *http.Client) (*ClientRegistrationResponse, error) {
|
||||
if registrationEndpoint == "" {
|
||||
return nil, fmt.Errorf("registration_endpoint is required")
|
||||
}
|
||||
|
||||
if c == nil {
|
||||
c = http.DefaultClient
|
||||
}
|
||||
|
||||
payload, err := json.Marshal(clientMeta)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal client metadata: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", registrationEndpoint, bytes.NewBuffer(payload))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create registration request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Accept", "application/json")
|
||||
|
||||
resp, err := c.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("registration request failed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read registration response body: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode == http.StatusCreated {
|
||||
var regResponse ClientRegistrationResponse
|
||||
if err := internaljson.Unmarshal(body, ®Response); err != nil {
|
||||
return nil, fmt.Errorf("failed to decode successful registration response: %w (%s)", err, string(body))
|
||||
}
|
||||
if regResponse.ClientID == "" {
|
||||
return nil, fmt.Errorf("registration response is missing required 'client_id' field")
|
||||
}
|
||||
// Validate URL fields to prevent XSS attacks (see #526).
|
||||
if err := validateClientRegistrationURLs(®Response.ClientRegistrationMetadata); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return ®Response, nil
|
||||
}
|
||||
|
||||
if resp.StatusCode == http.StatusBadRequest {
|
||||
var regError ClientRegistrationError
|
||||
if err := internaljson.Unmarshal(body, ®Error); err != nil {
|
||||
return nil, fmt.Errorf("failed to decode registration error response: %w (%s)", err, string(body))
|
||||
}
|
||||
return nil, ®Error
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf("registration failed with status %s: %s", resp.Status, string(body))
|
||||
}
|
||||
|
||||
// validateClientRegistrationURLs validates all URL fields in ClientRegistrationMetadata
|
||||
// to ensure they don't use dangerous schemes that could enable XSS attacks.
|
||||
func validateClientRegistrationURLs(meta *ClientRegistrationMetadata) error {
|
||||
// Validate redirect URIs
|
||||
for i, uri := range meta.RedirectURIs {
|
||||
if err := checkURLScheme(uri); err != nil {
|
||||
return fmt.Errorf("redirect_uris[%d]: %w", i, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Validate other URL fields
|
||||
urls := []struct {
|
||||
name string
|
||||
value string
|
||||
}{
|
||||
{"client_uri", meta.ClientURI},
|
||||
{"logo_uri", meta.LogoURI},
|
||||
{"tos_uri", meta.TOSURI},
|
||||
{"policy_uri", meta.PolicyURI},
|
||||
{"jwks_uri", meta.JWKSURI},
|
||||
}
|
||||
|
||||
for _, u := range urls {
|
||||
if err := checkURLScheme(u.value); err != nil {
|
||||
return fmt.Errorf("%s: %w", u.name, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
+96
@@ -0,0 +1,96 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// Package oauthex implements extensions to OAuth2.
|
||||
|
||||
//go:build mcp_go_client_oauth
|
||||
|
||||
package oauthex
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/internal/util"
|
||||
)
|
||||
|
||||
type httpStatusError struct {
|
||||
StatusCode int
|
||||
}
|
||||
|
||||
func (e *httpStatusError) Error() string {
|
||||
return fmt.Sprintf("bad status %d", e.StatusCode)
|
||||
}
|
||||
|
||||
// getJSON retrieves JSON and unmarshals JSON from the URL, as specified in both
|
||||
// RFC 9728 and RFC 8414.
|
||||
// It will not read more than limit bytes from the body.
|
||||
func getJSON[T any](ctx context.Context, c *http.Client, url string, limit int64) (*T, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if c == nil {
|
||||
c = http.DefaultClient
|
||||
}
|
||||
res, err := c.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer res.Body.Close()
|
||||
|
||||
if res.StatusCode != http.StatusOK {
|
||||
return nil, &httpStatusError{StatusCode: res.StatusCode}
|
||||
}
|
||||
ct := res.Header.Get("Content-Type")
|
||||
mediaType, _, err := mime.ParseMediaType(ct)
|
||||
if err != nil || mediaType != "application/json" {
|
||||
return nil, fmt.Errorf("bad content type %q", ct)
|
||||
}
|
||||
|
||||
var t T
|
||||
dec := json.NewDecoder(io.LimitReader(res.Body, limit))
|
||||
if err := dec.Decode(&t); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &t, nil
|
||||
}
|
||||
|
||||
// checkURLScheme ensures that its argument is a valid URL with a scheme
|
||||
// that prevents XSS attacks.
|
||||
// See #526.
|
||||
func checkURLScheme(u string) error {
|
||||
if u == "" {
|
||||
return nil
|
||||
}
|
||||
uu, err := url.Parse(u)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
scheme := strings.ToLower(uu.Scheme)
|
||||
if scheme == "javascript" || scheme == "data" || scheme == "vbscript" {
|
||||
return fmt.Errorf("URL has disallowed scheme %q", scheme)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func checkHTTPSOrLoopback(addr string) error {
|
||||
if addr == "" {
|
||||
return nil
|
||||
}
|
||||
u, err := url.Parse(addr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !util.IsLoopback(u.Host) && u.Scheme != "https" {
|
||||
return fmt.Errorf("URL %q does not use HTTPS or is not a loopback address", addr)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
+6
@@ -0,0 +1,6 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// Package oauthex implements extensions to OAuth2.
|
||||
package oauthex
|
||||
+279
@@ -0,0 +1,279 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// This file implements Protected Resource Metadata.
|
||||
// See https://www.rfc-editor.org/rfc/rfc9728.html.
|
||||
|
||||
//go:build mcp_go_client_oauth
|
||||
|
||||
package oauthex
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"path"
|
||||
"strings"
|
||||
"unicode"
|
||||
|
||||
"github.com/modelcontextprotocol/go-sdk/internal/util"
|
||||
)
|
||||
|
||||
const defaultProtectedResourceMetadataURI = "/.well-known/oauth-protected-resource"
|
||||
|
||||
// GetProtectedResourceMetadataFromID issues a GET request to retrieve protected resource
|
||||
// metadata from a resource server by its ID.
|
||||
// The resource ID is an HTTPS URL, typically with a host:port and possibly a path.
|
||||
// For example:
|
||||
//
|
||||
// https://example.com/server
|
||||
//
|
||||
// This function, following the spec (§3), inserts the default well-known path into the
|
||||
// URL. In our example, the result would be
|
||||
//
|
||||
// https://example.com/.well-known/oauth-protected-resource/server
|
||||
//
|
||||
// It then retrieves the metadata at that location using the given client (or the
|
||||
// default client if nil) and validates its resource field against resourceID.
|
||||
//
|
||||
// Deprecated: Use [GetProtectedResourceMetadata] instead. This function will be removed in v1.5.0.
|
||||
func GetProtectedResourceMetadataFromID(ctx context.Context, resourceID string, c *http.Client) (_ *ProtectedResourceMetadata, err error) {
|
||||
defer util.Wrapf(&err, "GetProtectedResourceMetadataFromID(%q)", resourceID)
|
||||
|
||||
u, err := url.Parse(resourceID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Insert well-known URI into URL.
|
||||
u.Path = path.Join(defaultProtectedResourceMetadataURI, u.Path)
|
||||
return GetProtectedResourceMetadata(ctx, u.String(), resourceID, c)
|
||||
}
|
||||
|
||||
// GetProtectedResourceMetadataFromHeader retrieves protected resource metadata
|
||||
// using information in the given header, using the given client (or the default
|
||||
// client if nil).
|
||||
// It issues a GET request to a URL discovered by parsing the WWW-Authenticate headers in the given request.
|
||||
// Per RFC 9728 section 3.3, it validates that the resource field of the resulting metadata
|
||||
// matches the serverURL (the URL that the client used to make the original request to the resource server).
|
||||
// If there is no metadata URL in the header, it returns nil, nil.
|
||||
//
|
||||
// Deprecated: Use [GetProtectedResourceMetadata] instead. This function will be removed in v1.5.0.
|
||||
func GetProtectedResourceMetadataFromHeader(ctx context.Context, serverURL string, header http.Header, c *http.Client) (_ *ProtectedResourceMetadata, err error) {
|
||||
headers := header[http.CanonicalHeaderKey("WWW-Authenticate")]
|
||||
if len(headers) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
cs, err := ParseWWWAuthenticate(headers)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
metadataURL := resourceMetadataURL(cs)
|
||||
if metadataURL == "" {
|
||||
return nil, nil
|
||||
}
|
||||
return GetProtectedResourceMetadata(ctx, metadataURL, serverURL, c)
|
||||
}
|
||||
|
||||
// resourceMetadataURL returns a resource metadata URL from the given "WWW-Authenticate" header challenges,
|
||||
// or the empty string if there is none.
|
||||
func resourceMetadataURL(cs []Challenge) string {
|
||||
for _, c := range cs {
|
||||
if u := c.Params["resource_metadata"]; u != "" {
|
||||
return u
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// GetProtectedResourceMetadataFromID issues a GET request to retrieve protected resource
|
||||
// metadata from a resource server.
|
||||
// The metadataURL is typically a URL with a host:port and possibly a path.
|
||||
// The resourceURL is the resource URI the metadataURL is for.
|
||||
// The following checks are performed:
|
||||
// - The metadataURL must use HTTPS or be a local address.
|
||||
// - The resource field of the resulting metadata must match the resourceURL.
|
||||
// - The authorization_servers field of the resulting metadata is checked for dangerous URL schemes.
|
||||
func GetProtectedResourceMetadata(ctx context.Context, metadataURL, resourceURL string, c *http.Client) (_ *ProtectedResourceMetadata, err error) {
|
||||
defer util.Wrapf(&err, "GetProtectedResourceMetadata(%q)", metadataURL)
|
||||
// Only allow HTTP for local addresses (testing or development purposes).
|
||||
if err := checkHTTPSOrLoopback(metadataURL); err != nil {
|
||||
return nil, fmt.Errorf("metadataURL: %v", err)
|
||||
}
|
||||
prm, err := getJSON[ProtectedResourceMetadata](ctx, c, metadataURL, 1<<20)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Validate the Resource field (see RFC 9728, section 3.3).
|
||||
if prm.Resource != resourceURL {
|
||||
return nil, fmt.Errorf("got metadata resource %q, want %q", prm.Resource, resourceURL)
|
||||
}
|
||||
// Validate the authorization server URLs to prevent XSS attacks (see #526).
|
||||
for i, u := range prm.AuthorizationServers {
|
||||
if err := checkURLScheme(u); err != nil {
|
||||
return nil, fmt.Errorf("authorization_servers[%d]: %v", i, err)
|
||||
}
|
||||
if err := checkHTTPSOrLoopback(u); err != nil {
|
||||
return nil, fmt.Errorf("authorization_servers[%d]: %v", i, err)
|
||||
}
|
||||
}
|
||||
return prm, nil
|
||||
}
|
||||
|
||||
// ParseWWWAuthenticate parses a WWW-Authenticate header string.
|
||||
// The header format is defined in RFC 9110, Section 11.6.1, and can contain
|
||||
// one or more challenges, separated by commas.
|
||||
// It returns a slice of challenges or an error if one of the headers is malformed.
|
||||
func ParseWWWAuthenticate(headers []string) ([]Challenge, error) {
|
||||
var challenges []Challenge
|
||||
for _, h := range headers {
|
||||
challengeStrings, err := splitChallenges(h)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, cs := range challengeStrings {
|
||||
if strings.TrimSpace(cs) == "" {
|
||||
continue
|
||||
}
|
||||
challenge, err := parseSingleChallenge(cs)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse challenge %q: %w", cs, err)
|
||||
}
|
||||
challenges = append(challenges, challenge)
|
||||
}
|
||||
}
|
||||
return challenges, nil
|
||||
}
|
||||
|
||||
// splitChallenges splits a header value containing one or more challenges.
|
||||
// It correctly handles commas within quoted strings and distinguishes between
|
||||
// commas separating auth-params and commas separating challenges.
|
||||
func splitChallenges(header string) ([]string, error) {
|
||||
var challenges []string
|
||||
inQuotes := false
|
||||
start := 0
|
||||
for i, r := range header {
|
||||
if r == '"' {
|
||||
if i > 0 && header[i-1] != '\\' {
|
||||
inQuotes = !inQuotes
|
||||
} else if i == 0 {
|
||||
// A challenge begins with an auth-scheme, which is a token, which cannot contain
|
||||
// a quote.
|
||||
return nil, errors.New(`challenge begins with '"'`)
|
||||
}
|
||||
} else if r == ',' && !inQuotes {
|
||||
// This is a potential challenge separator.
|
||||
// A new challenge does not start with `key=value`.
|
||||
// We check if the part after the comma looks like a parameter.
|
||||
lookahead := strings.TrimSpace(header[i+1:])
|
||||
eqPos := strings.Index(lookahead, "=")
|
||||
|
||||
isParam := false
|
||||
if eqPos > 0 {
|
||||
// Check if the part before '=' is a single token (no spaces).
|
||||
token := lookahead[:eqPos]
|
||||
if strings.IndexFunc(token, unicode.IsSpace) == -1 {
|
||||
isParam = true
|
||||
}
|
||||
}
|
||||
|
||||
if !isParam {
|
||||
// The part after the comma does not look like a parameter,
|
||||
// so this comma separates challenges.
|
||||
challenges = append(challenges, header[start:i])
|
||||
start = i + 1
|
||||
}
|
||||
}
|
||||
}
|
||||
// Add the last (or only) challenge to the list.
|
||||
challenges = append(challenges, header[start:])
|
||||
return challenges, nil
|
||||
}
|
||||
|
||||
// parseSingleChallenge parses a string containing exactly one challenge.
|
||||
// challenge = auth-scheme [ 1*SP ( token68 / #auth-param ) ]
|
||||
func parseSingleChallenge(s string) (Challenge, error) {
|
||||
s = strings.TrimSpace(s)
|
||||
if s == "" {
|
||||
return Challenge{}, errors.New("empty challenge string")
|
||||
}
|
||||
|
||||
scheme, paramsStr, found := strings.Cut(s, " ")
|
||||
c := Challenge{Scheme: strings.ToLower(scheme)}
|
||||
if !found {
|
||||
return c, nil
|
||||
}
|
||||
|
||||
params := make(map[string]string)
|
||||
|
||||
// Parse the key-value parameters.
|
||||
for paramsStr != "" {
|
||||
// Find the end of the parameter key.
|
||||
keyEnd := strings.Index(paramsStr, "=")
|
||||
if keyEnd <= 0 {
|
||||
return Challenge{}, fmt.Errorf("malformed auth parameter: expected key=value, but got %q", paramsStr)
|
||||
}
|
||||
key := strings.TrimSpace(paramsStr[:keyEnd])
|
||||
|
||||
// Move the string past the key and the '='.
|
||||
paramsStr = strings.TrimSpace(paramsStr[keyEnd+1:])
|
||||
|
||||
var value string
|
||||
if strings.HasPrefix(paramsStr, "\"") {
|
||||
// The value is a quoted string.
|
||||
paramsStr = paramsStr[1:] // Consume the opening quote.
|
||||
var valBuilder strings.Builder
|
||||
i := 0
|
||||
for ; i < len(paramsStr); i++ {
|
||||
// Handle escaped characters.
|
||||
if paramsStr[i] == '\\' && i+1 < len(paramsStr) {
|
||||
valBuilder.WriteByte(paramsStr[i+1])
|
||||
i++ // We've consumed two characters.
|
||||
} else if paramsStr[i] == '"' {
|
||||
// End of the quoted string.
|
||||
break
|
||||
} else {
|
||||
valBuilder.WriteByte(paramsStr[i])
|
||||
}
|
||||
}
|
||||
|
||||
// A quoted string must be terminated.
|
||||
if i == len(paramsStr) {
|
||||
return Challenge{}, fmt.Errorf("unterminated quoted string in auth parameter")
|
||||
}
|
||||
|
||||
value = valBuilder.String()
|
||||
// Move the string past the value and the closing quote.
|
||||
paramsStr = strings.TrimSpace(paramsStr[i+1:])
|
||||
} else {
|
||||
// The value is a token. It ends at the next comma or the end of the string.
|
||||
commaPos := strings.Index(paramsStr, ",")
|
||||
if commaPos == -1 {
|
||||
value = paramsStr
|
||||
paramsStr = ""
|
||||
} else {
|
||||
value = strings.TrimSpace(paramsStr[:commaPos])
|
||||
paramsStr = strings.TrimSpace(paramsStr[commaPos:]) // Keep comma for next check
|
||||
}
|
||||
}
|
||||
if value == "" {
|
||||
return Challenge{}, fmt.Errorf("no value for auth param %q", key)
|
||||
}
|
||||
|
||||
// Per RFC 9110, parameter keys are case-insensitive.
|
||||
params[strings.ToLower(key)] = value
|
||||
|
||||
// If there is a comma, consume it and continue to the next parameter.
|
||||
if strings.HasPrefix(paramsStr, ",") {
|
||||
paramsStr = strings.TrimSpace(paramsStr[1:])
|
||||
} else if paramsStr != "" {
|
||||
// If there's content but it's not a new parameter, the format is wrong.
|
||||
return Challenge{}, fmt.Errorf("malformed auth parameter: expected comma after value, but got %q", paramsStr)
|
||||
}
|
||||
}
|
||||
|
||||
// Per RFC 9110, the scheme is case-insensitive.
|
||||
return Challenge{Scheme: strings.ToLower(scheme), Params: params}, nil
|
||||
}
|
||||
+105
@@ -0,0 +1,105 @@
|
||||
// Copyright 2025 The Go MCP SDK Authors. All rights reserved.
|
||||
// Use of this source code is governed by an MIT-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// This file implements Protected Resource Metadata.
|
||||
// See https://www.rfc-editor.org/rfc/rfc9728.html.
|
||||
|
||||
// This is a temporary file to expose the required objects to the main package.
|
||||
|
||||
package oauthex
|
||||
|
||||
// ProtectedResourceMetadata is the metadata for an OAuth 2.0 protected resource,
|
||||
// as defined in section 2 of https://www.rfc-editor.org/rfc/rfc9728.html.
|
||||
//
|
||||
// The following features are not supported:
|
||||
// - additional keys (§2, last sentence)
|
||||
// - human-readable metadata (§2.1)
|
||||
// - signed metadata (§2.2)
|
||||
type ProtectedResourceMetadata struct {
|
||||
// Resource (resource) is the protected resource's resource identifier.
|
||||
// Required.
|
||||
Resource string `json:"resource"`
|
||||
|
||||
// AuthorizationServers (authorization_servers) is an optional slice containing a list of
|
||||
// OAuth authorization server issuer identifiers (as defined in RFC 8414) that can be
|
||||
// used with this protected resource.
|
||||
AuthorizationServers []string `json:"authorization_servers,omitempty"`
|
||||
|
||||
// JWKSURI (jwks_uri) is an optional URL of the protected resource's JSON Web Key (JWK) Set
|
||||
// document. This contains public keys belonging to the protected resource, such as
|
||||
// signing key(s) that the resource server uses to sign resource responses.
|
||||
JWKSURI string `json:"jwks_uri,omitempty"`
|
||||
|
||||
// ScopesSupported (scopes_supported) is a recommended slice containing a list of scope
|
||||
// values (as defined in RFC 6749) used in authorization requests to request access
|
||||
// to this protected resource.
|
||||
ScopesSupported []string `json:"scopes_supported,omitempty"`
|
||||
|
||||
// BearerMethodsSupported (bearer_methods_supported) is an optional slice containing
|
||||
// a list of the supported methods of sending an OAuth 2.0 bearer token to the
|
||||
// protected resource. Defined values are "header", "body", and "query".
|
||||
BearerMethodsSupported []string `json:"bearer_methods_supported,omitempty"`
|
||||
|
||||
// ResourceSigningAlgValuesSupported (resource_signing_alg_values_supported) is an optional
|
||||
// slice of JWS signing algorithms (alg values) supported by the protected
|
||||
// resource for signing resource responses.
|
||||
ResourceSigningAlgValuesSupported []string `json:"resource_signing_alg_values_supported,omitempty"`
|
||||
|
||||
// ResourceName (resource_name) is a human-readable name of the protected resource
|
||||
// intended for display to the end user. It is RECOMMENDED that this field be included.
|
||||
// This value may be internationalized.
|
||||
ResourceName string `json:"resource_name,omitempty"`
|
||||
|
||||
// ResourceDocumentation (resource_documentation) is an optional URL of a page containing
|
||||
// human-readable information for developers using the protected resource.
|
||||
// This value may be internationalized.
|
||||
ResourceDocumentation string `json:"resource_documentation,omitempty"`
|
||||
|
||||
// ResourcePolicyURI (resource_policy_uri) is an optional URL of a page containing
|
||||
// human-readable policy information on how a client can use the data provided.
|
||||
// This value may be internationalized.
|
||||
ResourcePolicyURI string `json:"resource_policy_uri,omitempty"`
|
||||
|
||||
// ResourceTOSURI (resource_tos_uri) is an optional URL of a page containing the protected
|
||||
// resource's human-readable terms of service. This value may be internationalized.
|
||||
ResourceTOSURI string `json:"resource_tos_uri,omitempty"`
|
||||
|
||||
// TLSClientCertificateBoundAccessTokens (tls_client_certificate_bound_access_tokens) is an
|
||||
// optional boolean indicating support for mutual-TLS client certificate-bound
|
||||
// access tokens (RFC 8705). Defaults to false if omitted.
|
||||
TLSClientCertificateBoundAccessTokens bool `json:"tls_client_certificate_bound_access_tokens,omitempty"`
|
||||
|
||||
// AuthorizationDetailsTypesSupported (authorization_details_types_supported) is an optional
|
||||
// slice of 'type' values supported by the resource server for the
|
||||
// 'authorization_details' parameter (RFC 9396).
|
||||
AuthorizationDetailsTypesSupported []string `json:"authorization_details_types_supported,omitempty"`
|
||||
|
||||
// DPOPSigningAlgValuesSupported (dpop_signing_alg_values_supported) is an optional
|
||||
// slice of JWS signing algorithms supported by the resource server for validating
|
||||
// DPoP proof JWTs (RFC 9449).
|
||||
DPOPSigningAlgValuesSupported []string `json:"dpop_signing_alg_values_supported,omitempty"`
|
||||
|
||||
// DPOPBoundAccessTokensRequired (dpop_bound_access_tokens_required) is an optional boolean
|
||||
// specifying whether the protected resource always requires the use of DPoP-bound
|
||||
// access tokens (RFC 9449). Defaults to false if omitted.
|
||||
DPOPBoundAccessTokensRequired bool `json:"dpop_bound_access_tokens_required,omitempty"`
|
||||
|
||||
// SignedMetadata (signed_metadata) is an optional JWT containing metadata parameters
|
||||
// about the protected resource as claims. If present, these values take precedence
|
||||
// over values conveyed in plain JSON.
|
||||
// TODO:implement.
|
||||
// Note that §2.2 says it's okay to ignore this.
|
||||
// SignedMetadata string `json:"signed_metadata,omitempty"`
|
||||
}
|
||||
|
||||
// Challenge represents a single authentication challenge from a WWW-Authenticate header.
|
||||
// As per RFC 9110, Section 11.6.1, a challenge consists of a scheme and optional parameters.
|
||||
type Challenge struct {
|
||||
// Scheme is the authentication scheme (e.g., "Bearer", "Basic").
|
||||
// It is case-insensitive. A parsed value will always be lower-case.
|
||||
Scheme string
|
||||
// Params is a map of authentication parameters.
|
||||
// Keys are case-insensitive. Parsed keys are always lower-case.
|
||||
Params map[string]string
|
||||
}
|
||||
Reference in New Issue
Block a user