diff --git a/wristband/register.go b/wristband/register.go index 38f2e47..0dda298 100644 --- a/wristband/register.go +++ b/wristband/register.go @@ -1,13 +1,332 @@ -// RFC 7591 Dynamic Client Registration (08-02-PLAN.md Task 2, D-02/D-05/ -// D-06/D-21). This is the Phase 8 Wave 2 RED placeholder: Register returns -// 501 so TestPhase8RedRegistration fails with the PHASE8_RED:registration -// sentinel while scripts/check-phase8-red.sh verifies the failure is -// fail-closed, matching the 08-01-PLAN.md Task 1 RED/GREEN pattern. +// RFC 7591 Dynamic Client Registration, ported from PHP +// OAuthRegisterController::register byte-for-byte including its validation +// order (08-CONTEXT.md D-02/D-05/D-06/D-21; canonical PHP source: +// OAuthRegisterController.php). package wristband -import "net/http" +import ( + "encoding/json" + "errors" + "net/http" + "net/url" + "strings" +) -// Register handles POST /oauth/mcp/register. -func (s *Server) Register(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusNotImplemented) +var supportedAuthMethods = map[string]bool{ + "none": true, + "client_secret_post": true, + "client_secret_basic": true, +} + +var supportedGrantTypes = []string{"authorization_code", "refresh_token"} + +// registerRequestBody is the RFC 7591 JSON request document. Only the +// fields PHP reads are decoded; unknown fields are ignored (PHP's +// $request->input() does the same). +type registerRequestBody struct { + RedirectURIs []string `json:"redirect_uris"` + TokenEndpointAuthMethod string `json:"token_endpoint_auth_method"` + GrantTypes []string `json:"grant_types"` + ResponseTypes []string `json:"response_types"` + ClientName string `json:"client_name"` +} + +type registerResponseBody struct { + ClientID string `json:"client_id"` + ClientIDIssuedAt int64 `json:"client_id_issued_at"` + ClientName string `json:"client_name"` + RedirectURIs []string `json:"redirect_uris"` + GrantTypes []string `json:"grant_types"` + ResponseTypes []string `json:"response_types"` + TokenEndpointAuthMethod string `json:"token_endpoint_auth_method"` + ClientSecret string `json:"client_secret,omitempty"` + ClientSecretExpiresAt *int64 `json:"client_secret_expires_at,omitempty"` +} + +type rfcErrorBody struct { + Error string `json:"error"` + ErrorDescription string `json:"error_description"` +} + +// Register handles POST /oauth/mcp/register (RFC 7591). D-02: JSON-only, +// PHP isJson() semantics ("contains /json"). D-21: the body is bounded to +// Options.RegisterMaxBodyBytes before decoding, and an oversized body gets +// the endpoint's ordinary invalid_client_metadata response, not a generic +// error. Sweep, the cap check and the create all run inside one transaction +// (T-08-DCR-FLOOD). +func (s *Server) Register(w http.ResponseWriter, r *http.Request) { + if !isJSONContentType(r.Header.Get("Content-Type")) { + writeRegisterError(w, http.StatusBadRequest, "invalid_client_metadata", "Request must be application/json") + return + } + if s.backend == nil { + w.WriteHeader(http.StatusInternalServerError) + return + } + + r.Body = http.MaxBytesReader(w, r.Body, s.opts.RegisterMaxBodyBytes) + var body registerRequestBody + dec := json.NewDecoder(r.Body) + if err := dec.Decode(&body); err != nil { + // D-21: an oversized (http.MaxBytesReader) or otherwise malformed + // JSON body both collapse to the endpoint-native error, never a + // generic 413/house body. + writeRegisterError(w, http.StatusBadRequest, "invalid_client_metadata", "Request must be application/json") + return + } + + uris, errMsg := parseRedirectURIs(body.RedirectURIs) + if errMsg != "" { + writeRegisterError(w, http.StatusBadRequest, "invalid_redirect_uri", errMsg) + return + } + + authMethod := body.TokenEndpointAuthMethod + if authMethod == "" { + authMethod = "none" + } + if !supportedAuthMethods[authMethod] { + writeRegisterError(w, http.StatusBadRequest, "invalid_client_metadata", "token_endpoint_auth_method is not supported") + return + } + + grants, errMsg := parseGrantTypes(body.GrantTypes) + if errMsg != "" { + writeRegisterError(w, http.StatusBadRequest, "invalid_client_metadata", errMsg) + return + } + + responseTypes, errMsg := parseResponseTypes(body.ResponseTypes) + if errMsg != "" { + writeRegisterError(w, http.StatusBadRequest, "invalid_client_metadata", errMsg) + return + } + + name := truncateRunes(strings.TrimSpace(stripControlChars(body.ClientName)), 255) + if name == "" { + name = "MCP client" + } + + clientID, err := s.randomBytes(16) + if err != nil { + w.WriteHeader(http.StatusInternalServerError) + return + } + var secret string + var secretHash *string + if authMethod != "none" { + secret, err = s.randomBytes(32) + if err != nil { + w.WriteHeader(http.StatusInternalServerError) + return + } + h := sha256Hex(secret) + secretHash = &h + } + + ip := clientIP(r) + rec := &ClientRecord{ + ClientID: clientID, + ClientSecretHash: secretHash, + ClientName: name, + RedirectURIs: uris, + GrantTypes: grants, + TokenEndpointAuthMethod: authMethod, + RegistrationIP: &ip, + } + + ctx := r.Context() + sweepBefore := s.now().Add(-s.opts.DCRUnconsentedSweepAge) + err = s.backend.WithinTx(ctx, func(tx Tx) error { + if err := tx.SweepUnconsented(ctx, sweepBefore); err != nil { + return err + } + return tx.CreateWithCap(ctx, rec, s.opts.DCRClientCap) + }) + if errors.Is(err, ErrClientCapReached) { + writeRegisterError(w, http.StatusBadRequest, "invalid_client_metadata", "Registration temporarily unavailable") + return + } + if err != nil { + w.WriteHeader(http.StatusInternalServerError) + return + } + + resp := registerResponseBody{ + ClientID: clientID, + ClientIDIssuedAt: rec.CreatedAt.Unix(), + ClientName: name, + RedirectURIs: uris, + GrantTypes: grants, + ResponseTypes: responseTypes, + TokenEndpointAuthMethod: authMethod, + } + if secret != "" { + resp.ClientSecret = secret + zero := int64(0) + resp.ClientSecretExpiresAt = &zero + } + writeExactJSON(w, http.StatusCreated, resp, map[string]string{"Cache-Control": "no-store"}) +} + +func writeRegisterError(w http.ResponseWriter, status int, code, description string) { + writeExactJSON(w, status, rfcErrorBody{Error: code, ErrorDescription: description}, map[string]string{"Cache-Control": "no-store"}) +} + +// isJSONContentType mirrors Laravel's Request::isJson(): the Content-Type +// header contains "/json" or "+json" anywhere (D-02). +func isJSONContentType(ct string) bool { + ct = strings.ToLower(ct) + return strings.Contains(ct, "/json") || strings.Contains(ct, "+json") +} + +// parseRedirectURIs ports OAuthRegisterController::parseRedirectUris: a +// required, non-empty array of unique URI strings, at most 5, each passing +// rejectRedirectURI. +func parseRedirectURIs(raw []string) ([]string, string) { + if len(raw) == 0 { + return nil, "redirect_uris is required" + } + seen := make(map[string]bool, len(raw)) + uris := make([]string, 0, len(raw)) + for _, entry := range raw { + if entry == "" { + return nil, "redirect_uris must be an array of URI strings" + } + if seen[entry] { + continue + } + seen[entry] = true + uris = append(uris, entry) + } + if len(uris) > 5 { + return nil, "A client may have at most 5 redirect URIs" + } + for _, u := range uris { + if reason := rejectRedirectURI(u); reason != "" { + return nil, reason + } + } + return uris, "" +} + +// rejectRedirectURI ports OAuthClient::rejectRedirectUri: at most 512 +// characters, a valid URL with scheme+host, https:// or loopback +// http://127.0.0.1 / http://localhost. Returns "" when the URI is accepted. +func rejectRedirectURI(uri string) string { + if len(uri) > 512 { + return "Each redirect URI must be at most 512 characters." + } + parsed, err := url.Parse(uri) + if err != nil || parsed.Scheme == "" || parsed.Hostname() == "" { + return "Redirect URI is not a valid URL: " + uri + } + scheme := strings.ToLower(parsed.Scheme) + host := strings.ToLower(parsed.Hostname()) + if scheme == "https" { + return "" + } + if scheme == "http" && (host == "127.0.0.1" || host == "localhost") { + return "" + } + return "Redirect URI must be https:// or loopback http://127.0.0.1 / http://localhost: " + uri +} + +// parseGrantTypes ports OAuthRegisterController::parseGrantTypes: nil means +// the default (both supported grants); otherwise a non-empty array of +// supported values that must include authorization_code. +func parseGrantTypes(raw []string) ([]string, string) { + if raw == nil { + return append([]string(nil), supportedGrantTypes...), "" + } + if len(raw) == 0 { + return nil, "grant_types is invalid" + } + seen := make(map[string]bool, len(raw)) + grants := make([]string, 0, len(raw)) + hasAuthCode := false + for _, entry := range raw { + if !isSupportedGrantType(entry) { + return nil, "grant_types contains an unsupported value" + } + if seen[entry] { + continue + } + seen[entry] = true + grants = append(grants, entry) + if entry == "authorization_code" { + hasAuthCode = true + } + } + if !hasAuthCode { + return nil, "grant_types must include authorization_code" + } + return grants, "" +} + +func isSupportedGrantType(v string) bool { + for _, g := range supportedGrantTypes { + if g == v { + return true + } + } + return false +} + +// parseResponseTypes ports OAuthRegisterController::parseResponseTypes: +// nil means the default ["code"]; otherwise every entry must be "code". +func parseResponseTypes(raw []string) ([]string, string) { + if raw == nil { + return []string{"code"}, "" + } + if len(raw) == 0 { + return nil, "response_types is invalid" + } + for _, entry := range raw { + if entry != "code" { + return nil, "response_types must be code" + } + } + return []string{"code"}, "" +} + +// stripControlChars removes C0 control characters and DEL, mirroring the +// PHP preg_replace('/[\x00-\x1F\x7F]/u', ”, ...) regex ConnectedApp/Consent +// controllers apply to client_name at display time (08-CONTEXT.md +// discretion note). Registration applies the same filter at capture time so +// no untrusted control byte from an issuer's declared name is ever stored. +func stripControlChars(s string) string { + return strings.Map(func(r rune) rune { + if r <= 0x1F || r == 0x7F { + return -1 + } + return r + }, s) +} + +// truncateRunes returns s truncated to at most n runes (PHP mb_substr parity). +func truncateRunes(s string, n int) string { + r := []rune(s) + if len(r) <= n { + return s + } + return string(r[:n]) +} + +// clientIP mirrors PHP's substr((string) $request->ip(), 0, 45): the +// direct remote address (no trusted-proxy chain here; DCR is unauthenticated +// and app-specific proxy trust stays out of the app-agnostic framework). +func clientIP(r *http.Request) string { + host := r.RemoteAddr + if idx := strings.LastIndex(host, ":"); idx != -1 && !strings.Contains(host, "]") { + host = host[:idx] + } else if strings.HasPrefix(host, "[") { + if end := strings.Index(host, "]"); end != -1 { + host = host[1:end] + } + } + if len(host) > 45 { + host = host[:45] + } + return host }