fix(go.sum): update ResolveSpec dependency to v1.0.87
CI / build-and-test (push) Failing after 1s
Release / release (push) Failing after 19m26s

This commit is contained in:
Hein
2026-06-23 13:17:16 +02:00
parent 0227912325
commit 1adf50e3db
2436 changed files with 1078758 additions and 114 deletions
+216
View File
@@ -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
View File
@@ -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
}
})
}
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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 }
@@ -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
}
@@ -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
}
@@ -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
View File
@@ -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
}
}
@@ -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()
}
@@ -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
}
@@ -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
View File
@@ -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
View File
@@ -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)
}
}
@@ -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
View File
@@ -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
)
File diff suppressed because it is too large Load Diff
+108
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
}
File diff suppressed because it is too large Load Diff
+39
View File
@@ -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
View File
@@ -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
View File
@@ -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)
}
File diff suppressed because it is too large Load Diff
+29
View File
@@ -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
View File
@@ -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&notification != 0 && req.IsCall() {
return methodInfo{}, fmt.Errorf("%w: unexpected id for %q", jsonrpc2.ErrInvalidRequest, req.Method)
}
if info.flags&notification == 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
View File
@@ -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
}
File diff suppressed because it is too large Load Diff
+226
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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, &params); 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
View File
@@ -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
View File
@@ -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
View File
@@ -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, &regResponse); 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(&regResponse.ClientRegistrationMetadata); err != nil {
return nil, err
}
return &regResponse, nil
}
if resp.StatusCode == http.StatusBadRequest {
var regError ClientRegistrationError
if err := internaljson.Unmarshal(body, &regError); err != nil {
return nil, fmt.Errorf("failed to decode registration error response: %w (%s)", err, string(body))
}
return nil, &regError
}
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
View File
@@ -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
View File
@@ -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
View File
@@ -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
}
@@ -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
}