feat(08-03): implement exact authorize validation and pending creation in wristband
- Server.Authorize ports OAuthAuthorizeController::authorize's exact validation order: usable client, exact redirect, response_type=code, code_challenge_method=S256, challenge length, scope parsing/ceiling truncation, resource check, then opaque pending-request creation - Unknown client/unregistered redirect are local text/plain 400s with no Location; every later failure is an ordered RFC3986 redirect with error/error_description/iss[/state], built via a dedicated encoder (never url.Values.Encode, which sorts keys and space-encodes as '+') - Options gains Resource and PendingRequestTTL (both PHP-parity defaults) so authorize's resource check and 600s pending expiry are configurable
This commit is contained in:
@@ -13,13 +13,258 @@
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Authorize handles GET /oauth/mcp/authorize. It is not yet implemented
|
||||
// (Wave 3 Task 1 RED anchor, 08-03-PLAN.md); TestPhase8RedAuthorize and
|
||||
// TestPhase8RedAuthorizeApp fail against this stub until Task 2's GREEN
|
||||
// commit.
|
||||
func (s *Server) Authorize(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusNotImplemented)
|
||||
// authorizeAllowedScopes mirrors PHP OAuthAuthorizeController::ALLOWED_SCOPES.
|
||||
var authorizeAllowedScopes = map[string]bool{
|
||||
"read": true, "write": true, "ai": true, "offline_access": true,
|
||||
}
|
||||
|
||||
// Authorize handles GET /oauth/mcp/authorize (D-09: raw route, no
|
||||
// middleware). It validates the client and exact redirect before any
|
||||
// redirect response, enforces S256 PKCE syntax, the client's scope ceiling,
|
||||
// and the RFC 8707 resource check, then persists an opaque pending request
|
||||
// and redirects to the app's /connect handoff.
|
||||
func (s *Server) Authorize(w http.ResponseWriter, r *http.Request) {
|
||||
if s.backend == nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
ctx := r.Context()
|
||||
q := r.URL.Query()
|
||||
|
||||
clientID := q.Get("client_id")
|
||||
var client *ClientRecord
|
||||
if clientID != "" {
|
||||
err := s.backend.WithinTx(ctx, func(tx Tx) error {
|
||||
c, err := tx.ByClientID(ctx, clientID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
client = c
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
}
|
||||
if client == nil || client.RevokedAt != nil {
|
||||
writeAuthorizeLocalError(w, "Unknown client.")
|
||||
return
|
||||
}
|
||||
|
||||
redirectURI := q.Get("redirect_uri")
|
||||
if redirectURI == "" || !stringSliceContains(client.RedirectURIs, redirectURI) {
|
||||
writeAuthorizeLocalError(w, "Unregistered redirect URI.")
|
||||
return
|
||||
}
|
||||
|
||||
var state *string
|
||||
if v := q.Get("state"); v != "" {
|
||||
state = &v
|
||||
}
|
||||
|
||||
if q.Get("response_type") != "code" {
|
||||
s.authorizeErrorRedirect(w, redirectURI, "unsupported_response_type", "response_type must be code", state)
|
||||
return
|
||||
}
|
||||
|
||||
if q.Get("code_challenge_method") != "S256" {
|
||||
s.authorizeErrorRedirect(w, redirectURI, "invalid_request", "code_challenge_method must be S256", state)
|
||||
return
|
||||
}
|
||||
|
||||
challenge := q.Get("code_challenge")
|
||||
if l := len(challenge); l < 43 || l > 128 {
|
||||
s.authorizeErrorRedirect(w, redirectURI, "invalid_request", "code_challenge is required", state)
|
||||
return
|
||||
}
|
||||
|
||||
scopes, scopeErr := parseAuthorizeScopes(q.Get("scope"))
|
||||
if scopeErr != "" {
|
||||
s.authorizeErrorRedirect(w, redirectURI, "invalid_scope", scopeErr, state)
|
||||
return
|
||||
}
|
||||
|
||||
// Phase 07.13 (D-10/D-11/D-12 in the PHP source comment): a nil ceiling
|
||||
// is a no-op; offline_access always survives truncation because it is
|
||||
// peeled into a boolean below, not a data scope; truncating here (not
|
||||
// rejecting) means an over-broad request silently narrows instead of
|
||||
// failing, except when nothing data-bearing survives.
|
||||
if client.ScopeCeiling != nil {
|
||||
kept := make([]string, 0, len(scopes))
|
||||
for _, sc := range scopes {
|
||||
if sc == "offline_access" || stringSliceContains(client.ScopeCeiling, sc) {
|
||||
kept = append(kept, sc)
|
||||
}
|
||||
}
|
||||
dataScopes := make([]string, 0, len(kept))
|
||||
for _, sc := range kept {
|
||||
if sc != "offline_access" {
|
||||
dataScopes = append(dataScopes, sc)
|
||||
}
|
||||
}
|
||||
if len(dataScopes) == 0 {
|
||||
s.authorizeErrorRedirect(w, redirectURI, "invalid_scope", "requested scope is outside this client's ceiling", state)
|
||||
return
|
||||
}
|
||||
scopes = kept
|
||||
}
|
||||
|
||||
var resource *string
|
||||
if v := q.Get("resource"); v != "" {
|
||||
resource = &v
|
||||
}
|
||||
if resource != nil && *resource != s.opts.Resource {
|
||||
s.authorizeErrorRedirect(w, redirectURI, "invalid_target", "resource does not match this server", state)
|
||||
return
|
||||
}
|
||||
|
||||
// createPendingRequest (PHP OAuthCodeManager): peel offline_access into
|
||||
// a boolean flag; the persisted scopes list never contains it.
|
||||
offline := false
|
||||
dataScopes := make([]string, 0, len(scopes))
|
||||
for _, sc := range scopes {
|
||||
if sc == "offline_access" {
|
||||
offline = true
|
||||
continue
|
||||
}
|
||||
dataScopes = append(dataScopes, sc)
|
||||
}
|
||||
|
||||
requestID, err := s.randomBytes(32)
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
rec := &AuthCodeRecord{
|
||||
RequestID: &requestID,
|
||||
ClientID: client.ClientID,
|
||||
RedirectURI: redirectURI,
|
||||
Scopes: dataScopes,
|
||||
CodeChallenge: challenge,
|
||||
CodeChallengeMethod: "S256",
|
||||
Resource: resource,
|
||||
State: state,
|
||||
ExpiresAt: s.now().Add(s.opts.PendingRequestTTL),
|
||||
OfflineAccess: offline,
|
||||
}
|
||||
err = s.backend.WithinTx(ctx, func(tx Tx) error {
|
||||
return tx.CreatePending(ctx, rec)
|
||||
})
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
spa := s.opts.Issuer + "/connect?request=" + rfc3986Escape(requestID)
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.Header().Set("Location", spa)
|
||||
w.WriteHeader(http.StatusFound)
|
||||
}
|
||||
|
||||
// writeAuthorizeLocalError writes the PHP localError() response: a bare
|
||||
// text/plain 400 with no Location and no house envelope
|
||||
// (T-08-OPEN-REDIRECT).
|
||||
func writeAuthorizeLocalError(w http.ResponseWriter, message string) {
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.Header().Set("Content-Type", "text/plain; charset=UTF-8")
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
_, _ = w.Write([]byte(message))
|
||||
}
|
||||
|
||||
// authorizeErrorRedirect ports OAuthAuthorizeController::errorRedirect: an
|
||||
// ordered error/error_description/iss[/state] query appended to the
|
||||
// already-trusted redirectURI via RFC 3986 encoding (D-04/Pattern 3).
|
||||
func (s *Server) authorizeErrorRedirect(w http.ResponseWriter, redirectURI, errCode, description string, state *string) {
|
||||
pairs := [][2]string{
|
||||
{"error", errCode},
|
||||
{"error_description", description},
|
||||
{"iss", s.opts.Issuer},
|
||||
}
|
||||
if state != nil {
|
||||
pairs = append(pairs, [2]string{"state", *state})
|
||||
}
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.Header().Set("Location", appendOrderedQuery(redirectURI, pairs))
|
||||
w.WriteHeader(http.StatusFound)
|
||||
}
|
||||
|
||||
// parseAuthorizeScopes ports OAuthAuthorizeController::parseScopes. An
|
||||
// absent or empty scope query value defaults to ["read"]; otherwise every
|
||||
// whitespace-separated token must be one of authorizeAllowedScopes.
|
||||
func parseAuthorizeScopes(raw string) ([]string, string) {
|
||||
parts := strings.Fields(strings.TrimSpace(raw))
|
||||
if len(parts) == 0 {
|
||||
return []string{"read"}, ""
|
||||
}
|
||||
for _, p := range parts {
|
||||
if !authorizeAllowedScopes[p] {
|
||||
return nil, "scope contains an unsupported value"
|
||||
}
|
||||
}
|
||||
return parts, ""
|
||||
}
|
||||
|
||||
// stringSliceContains reports whether ss contains v.
|
||||
func stringSliceContains(ss []string, v string) bool {
|
||||
for _, s := range ss {
|
||||
if s == v {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// rfc3986Escape percent-encodes s per RFC 3986 (PHP rawurlencode/
|
||||
// PHP_QUERY_RFC3986): unreserved characters A-Z a-z 0-9 - _ . ~ pass
|
||||
// through unescaped; everything else, including the space byte, becomes an
|
||||
// uppercase %XX triplet. This intentionally differs from
|
||||
// net/url.QueryEscape (form-style '+' for space, different reserved-char
|
||||
// handling), which would corrupt the exact PHP redirect-byte contract
|
||||
// (08-RESEARCH.md Pattern 3/Pitfall 5).
|
||||
func rfc3986Escape(s string) string {
|
||||
var b strings.Builder
|
||||
b.Grow(len(s))
|
||||
for i := 0; i < len(s); i++ {
|
||||
c := s[i]
|
||||
if isRFC3986Unreserved(c) {
|
||||
b.WriteByte(c)
|
||||
} else {
|
||||
fmt.Fprintf(&b, "%%%02X", c)
|
||||
}
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func isRFC3986Unreserved(c byte) bool {
|
||||
return (c >= 'A' && c <= 'Z') || (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') ||
|
||||
c == '-' || c == '_' || c == '.' || c == '~'
|
||||
}
|
||||
|
||||
// buildOrderedQuery RFC3986-encodes and joins pairs in the given order
|
||||
// (never net/url.Values.Encode, which sorts keys and space-encodes as '+').
|
||||
func buildOrderedQuery(pairs [][2]string) string {
|
||||
parts := make([]string, len(pairs))
|
||||
for i, p := range pairs {
|
||||
parts[i] = rfc3986Escape(p[0]) + "=" + rfc3986Escape(p[1])
|
||||
}
|
||||
return strings.Join(parts, "&")
|
||||
}
|
||||
|
||||
// appendOrderedQuery ports OAuthAuthorizeController::appendQuery: appends an
|
||||
// ordered RFC 3986 query to uri, using '&' when uri already carries a query
|
||||
// string and '?' otherwise.
|
||||
func appendOrderedQuery(uri string, pairs [][2]string) string {
|
||||
sep := "?"
|
||||
if strings.Contains(uri, "?") {
|
||||
sep = "&"
|
||||
}
|
||||
return uri + sep + buildOrderedQuery(pairs)
|
||||
}
|
||||
|
||||
@@ -117,3 +117,614 @@ func TestPhase8RedAuthorize(t *testing.T) {
|
||||
t.Fatalf("PHASE8_RED:authorize: Cache-Control = %q, want \"no-store\"", cc)
|
||||
}
|
||||
}
|
||||
|
||||
const authorizeTestRedirect = "https://chatgpt.com/connector/oauth/cb"
|
||||
|
||||
func newAuthorizeTestServer() (*Server, *memoryBackend) {
|
||||
backend := newMemoryBackend()
|
||||
return newTestServer(backend), backend
|
||||
}
|
||||
|
||||
func pendingCount(backend *memoryBackend) int {
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
return len(backend.codes)
|
||||
}
|
||||
|
||||
func TestAuthorizeUnknownClientReturnsLocal400NoLocation(t *testing.T) {
|
||||
srv, _ := newAuthorizeTestServer()
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "does-not-exist",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want %d (body=%s)", rec.Code, http.StatusBadRequest, rec.Body.String())
|
||||
}
|
||||
if loc := rec.Header().Get("Location"); loc != "" {
|
||||
t.Fatalf("Location = %q, want none", loc)
|
||||
}
|
||||
if ct := rec.Header().Get("Content-Type"); ct != "text/plain; charset=UTF-8" {
|
||||
t.Fatalf("Content-Type = %q, want text/plain; charset=UTF-8", ct)
|
||||
}
|
||||
if body := rec.Body.String(); body != "Unknown client." {
|
||||
t.Fatalf("body = %q, want %q", body, "Unknown client.")
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-store" {
|
||||
t.Fatalf("Cache-Control = %q, want no-store", cc)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeRevokedClientIsUnknown(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
rec := insertAuthorizeTestClient(backend, "cli-revoked", []string{authorizeTestRedirect}, nil)
|
||||
now := time.Now()
|
||||
rec.RevokedAt = &now
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-revoked",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
})
|
||||
w := httptest.NewRecorder()
|
||||
srv.Authorize(w, req)
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want %d", w.Code, http.StatusBadRequest)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeUnregisteredRedirectURIReturnsLocal400NoLocation(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-redirect", []string{authorizeTestRedirect}, nil)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-redirect",
|
||||
"redirect_uri": "https://evil.test/cb",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusBadRequest)
|
||||
}
|
||||
if loc := rec.Header().Get("Location"); loc != "" {
|
||||
t.Fatalf("Location = %q, want none", loc)
|
||||
}
|
||||
if strings.Contains(rec.Body.String(), "https://evil.test/cb") {
|
||||
t.Fatal("body echoes the untrusted redirect URI")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeTrailingSlashMismatchIsUnregistered(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-slash", []string{authorizeTestRedirect}, nil)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-slash",
|
||||
"redirect_uri": authorizeTestRedirect + "/",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusBadRequest)
|
||||
}
|
||||
if loc := rec.Header().Get("Location"); loc != "" {
|
||||
t.Fatalf("Location = %q, want none", loc)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeUnsupportedResponseTypeRedirectsWithIssAndState(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-resptype", []string{authorizeTestRedirect}, nil)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-resptype",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "token",
|
||||
"code_challenge": strings.Repeat("b", 43),
|
||||
"code_challenge_method": "S256",
|
||||
"state": "iss-state",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusFound)
|
||||
}
|
||||
loc := rec.Header().Get("Location")
|
||||
if !strings.HasPrefix(loc, authorizeTestRedirect) {
|
||||
t.Fatalf("Location = %q, want prefix %q", loc, authorizeTestRedirect)
|
||||
}
|
||||
q := queryOf(t, loc)
|
||||
if q["error"] != "unsupported_response_type" {
|
||||
t.Fatalf("error = %q, want unsupported_response_type", q["error"])
|
||||
}
|
||||
if q["iss"] != "https://plytarium.com" {
|
||||
t.Fatalf("iss = %q, want https://plytarium.com", q["iss"])
|
||||
}
|
||||
if q["state"] != "iss-state" {
|
||||
t.Fatalf("state = %q, want iss-state", q["state"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestPKCEChallengeMethodMustBeS256(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-plain", []string{authorizeTestRedirect}, nil)
|
||||
state := "state-plain"
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-plain",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": strings.Repeat("a", 43),
|
||||
"code_challenge_method": "plain",
|
||||
"state": state,
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusFound)
|
||||
}
|
||||
q := queryOf(t, rec.Header().Get("Location"))
|
||||
if q["error"] != "invalid_request" {
|
||||
t.Fatalf("error = %q, want invalid_request", q["error"])
|
||||
}
|
||||
if q["state"] != state {
|
||||
t.Fatalf("state = %q, want %q", q["state"], state)
|
||||
}
|
||||
if q["iss"] != "https://plytarium.com" {
|
||||
t.Fatalf("iss = %q, want https://plytarium.com", q["iss"])
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-store" {
|
||||
t.Fatalf("Cache-Control = %q, want no-store", cc)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPKCEChallengeLengthBounds(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
challenge string
|
||||
}{
|
||||
{"too-short", strings.Repeat("a", 42)},
|
||||
{"too-long", strings.Repeat("a", 129)},
|
||||
{"empty", ""},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-len-"+tc.name, []string{authorizeTestRedirect}, nil)
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-len-" + tc.name,
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": tc.challenge,
|
||||
"code_challenge_method": "S256",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusFound)
|
||||
}
|
||||
q := queryOf(t, rec.Header().Get("Location"))
|
||||
if q["error"] != "invalid_request" {
|
||||
t.Fatalf("error = %q, want invalid_request", q["error"])
|
||||
}
|
||||
if q["error_description"] != "code_challenge is required" {
|
||||
t.Fatalf("error_description = %q, want %q", q["error_description"], "code_challenge is required")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeValidRequestRedirectsToConnectWithOpaqueHandleOnly(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-valid", []string{authorizeTestRedirect}, nil)
|
||||
_, challenge := s256Pair(t)
|
||||
state := "must-not-appear-on-our-url"
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-valid",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "read write",
|
||||
"state": state,
|
||||
"resource": "https://mcp.plytarium.com/mcp",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want %d (body=%s)", rec.Code, http.StatusFound, rec.Body.String())
|
||||
}
|
||||
loc := rec.Header().Get("Location")
|
||||
if !strings.HasPrefix(loc, "https://plytarium.com/connect?") {
|
||||
t.Fatalf("Location = %q, want prefix https://plytarium.com/connect?", loc)
|
||||
}
|
||||
q := queryOf(t, loc)
|
||||
if q["request"] == "" {
|
||||
t.Fatal("Location missing non-empty request param")
|
||||
}
|
||||
if _, has := q["code"]; has {
|
||||
t.Fatal("Location leaks a code")
|
||||
}
|
||||
if _, has := q["state"]; has {
|
||||
t.Fatal("Location leaks state")
|
||||
}
|
||||
if _, has := q["client_secret"]; has {
|
||||
t.Fatal("Location leaks client_secret")
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-store" {
|
||||
t.Fatalf("Cache-Control = %q, want no-store", cc)
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
if len(backend.codes) != 1 {
|
||||
t.Fatalf("pending rows = %d, want 1", len(backend.codes))
|
||||
}
|
||||
pending := backend.codes[0]
|
||||
if pending.RequestID == nil || *pending.RequestID != q["request"] {
|
||||
t.Fatalf("pending.RequestID = %v, want %q", pending.RequestID, q["request"])
|
||||
}
|
||||
if pending.CodeHash != nil {
|
||||
t.Fatal("pending row already has a code hash")
|
||||
}
|
||||
if pending.UserID != nil {
|
||||
t.Fatal("pending row already has a user id")
|
||||
}
|
||||
if pending.State == nil || *pending.State != state {
|
||||
t.Fatalf("pending.State = %v, want %q", pending.State, state)
|
||||
}
|
||||
if pending.ExpiresAt.Before(time.Now().Add(500*time.Second)) || pending.ExpiresAt.After(time.Now().Add(700*time.Second)) {
|
||||
t.Fatalf("pending.ExpiresAt = %v, want ~600s from now", pending.ExpiresAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeIssIsPresentOnEveryErrorRedirect(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-iss", []string{authorizeTestRedirect}, nil)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-iss",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "token",
|
||||
"code_challenge": strings.Repeat("b", 43),
|
||||
"code_challenge_method": "S256",
|
||||
"state": "iss-state",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
q := queryOf(t, rec.Header().Get("Location"))
|
||||
if q["iss"] != "https://plytarium.com" {
|
||||
t.Fatalf("iss = %q, want https://plytarium.com", q["iss"])
|
||||
}
|
||||
if q["state"] != "iss-state" {
|
||||
t.Fatalf("state = %q, want iss-state", q["state"])
|
||||
}
|
||||
if q["error"] != "unsupported_response_type" {
|
||||
t.Fatalf("error = %q, want unsupported_response_type", q["error"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeScopeDefaultsToRead(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-default-scope", []string{authorizeTestRedirect}, nil)
|
||||
_, challenge := s256Pair(t)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-default-scope",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusFound)
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
if len(backend.codes) != 1 {
|
||||
t.Fatalf("pending rows = %d, want 1", len(backend.codes))
|
||||
}
|
||||
if got := backend.codes[0].Scopes; len(got) != 1 || got[0] != "read" {
|
||||
t.Fatalf("scopes = %v, want [read]", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeScopeInvalidValueRedirectsInvalidScope(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-bad-scope", []string{authorizeTestRedirect}, nil)
|
||||
_, challenge := s256Pair(t)
|
||||
before := pendingCount(backend)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-bad-scope",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "read bogus",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusFound)
|
||||
}
|
||||
q := queryOf(t, rec.Header().Get("Location"))
|
||||
if q["error"] != "invalid_scope" {
|
||||
t.Fatalf("error = %q, want invalid_scope", q["error"])
|
||||
}
|
||||
if q["error_description"] != "scope contains an unsupported value" {
|
||||
t.Fatalf("error_description = %q", q["error_description"])
|
||||
}
|
||||
if got := pendingCount(backend); got != before {
|
||||
t.Fatalf("pending rows = %d, want unchanged %d", got, before)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeNullCeilingPreservesAiScope(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-null-ceiling", []string{authorizeTestRedirect}, nil)
|
||||
_, challenge := s256Pair(t)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-null-ceiling",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "read write ai",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
got := backend.codes[len(backend.codes)-1].Scopes
|
||||
want := []string{"read", "write", "ai"}
|
||||
if len(got) != len(want) {
|
||||
t.Fatalf("scopes = %v, want %v", got, want)
|
||||
}
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("scopes = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeCeilingTruncatesAiWithoutError(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-ceiling-trunc", []string{authorizeTestRedirect}, []string{"read", "write"})
|
||||
_, challenge := s256Pair(t)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-ceiling-trunc",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "read write ai",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
got := backend.codes[len(backend.codes)-1].Scopes
|
||||
if stringSliceContains(got, "ai") {
|
||||
t.Fatalf("scopes = %v, still contains ai", got)
|
||||
}
|
||||
want := []string{"read", "write"}
|
||||
if len(got) != len(want) || got[0] != want[0] || got[1] != want[1] {
|
||||
t.Fatalf("scopes = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeCeilingPreservesOfflineAccessFlag(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-ceiling-offline", []string{authorizeTestRedirect}, []string{"read", "write"})
|
||||
_, challenge := s256Pair(t)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-ceiling-offline",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "read write ai offline_access",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
pending := backend.codes[len(backend.codes)-1]
|
||||
want := []string{"read", "write"}
|
||||
if len(pending.Scopes) != len(want) || pending.Scopes[0] != want[0] || pending.Scopes[1] != want[1] {
|
||||
t.Fatalf("scopes = %v, want %v", pending.Scopes, want)
|
||||
}
|
||||
if !pending.OfflineAccess {
|
||||
t.Fatal("OfflineAccess = false, want true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeCeilingRejectsAiOnlyAsInvalidScopeOnTrustedRedirect(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-ceiling-ai-only", []string{authorizeTestRedirect}, []string{"read", "write"})
|
||||
_, challenge := s256Pair(t)
|
||||
state := "state-ai-only"
|
||||
before := pendingCount(backend)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-ceiling-ai-only",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "ai",
|
||||
"state": state,
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
loc := rec.Header().Get("Location")
|
||||
if !strings.HasPrefix(loc, authorizeTestRedirect) {
|
||||
t.Fatalf("Location = %q, want prefix %q", loc, authorizeTestRedirect)
|
||||
}
|
||||
q := queryOf(t, loc)
|
||||
if q["error"] != "invalid_scope" {
|
||||
t.Fatalf("error = %q, want invalid_scope", q["error"])
|
||||
}
|
||||
if q["state"] != state {
|
||||
t.Fatalf("state = %q, want %q", q["state"], state)
|
||||
}
|
||||
if q["iss"] != "https://plytarium.com" {
|
||||
t.Fatalf("iss = %q, want https://plytarium.com", q["iss"])
|
||||
}
|
||||
if got := pendingCount(backend); got != before {
|
||||
t.Fatalf("pending rows = %d, want unchanged %d", got, before)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeCeilingRejectsAiPlusOfflineAccessAsInvalidScope(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-ceiling-ai-offline", []string{authorizeTestRedirect}, []string{"read", "write"})
|
||||
_, challenge := s256Pair(t)
|
||||
state := "state-ai-offline"
|
||||
before := pendingCount(backend)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-ceiling-ai-offline",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "ai offline_access",
|
||||
"state": state,
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
q := queryOf(t, rec.Header().Get("Location"))
|
||||
if q["error"] != "invalid_scope" {
|
||||
t.Fatalf("error = %q, want invalid_scope", q["error"])
|
||||
}
|
||||
if q["state"] != state {
|
||||
t.Fatalf("state = %q, want %q", q["state"], state)
|
||||
}
|
||||
if got := pendingCount(backend); got != before {
|
||||
t.Fatalf("pending rows = %d, want unchanged %d", got, before)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeResourceMismatchRedirectsInvalidTarget(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-resource", []string{authorizeTestRedirect}, nil)
|
||||
_, challenge := s256Pair(t)
|
||||
before := pendingCount(backend)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-resource",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"resource": "https://wrong.example.test/mcp",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
q := queryOf(t, rec.Header().Get("Location"))
|
||||
if q["error"] != "invalid_target" {
|
||||
t.Fatalf("error = %q, want invalid_target", q["error"])
|
||||
}
|
||||
if q["error_description"] != "resource does not match this server" {
|
||||
t.Fatalf("error_description = %q", q["error_description"])
|
||||
}
|
||||
if got := pendingCount(backend); got != before {
|
||||
t.Fatalf("pending rows = %d, want unchanged %d", got, before)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeResourceOmittedIsAccepted(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-resource-omitted", []string{authorizeTestRedirect}, nil)
|
||||
_, challenge := s256Pair(t)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-resource-omitted",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
// TestOrderedRedirectQueryEncodingIsRFC3986 proves the redirect query
|
||||
// encoder differs from net/url.Values.Encode exactly where PHP
|
||||
// http_build_query(..., PHP_QUERY_RFC3986) does: spaces become %20 (not
|
||||
// '+'), reserved characters are percent-encoded, and key order is
|
||||
// preserved (not sorted) — Pitfall 5.
|
||||
func TestOrderedRedirectQueryEncodingIsRFC3986(t *testing.T) {
|
||||
got := buildOrderedQuery([][2]string{
|
||||
{"error", "invalid_scope"},
|
||||
{"error_description", "requested scope is outside this client's ceiling"},
|
||||
{"iss", "https://plytarium.com"},
|
||||
{"state", "a b&c=d"},
|
||||
})
|
||||
want := "error=invalid_scope&error_description=requested%20scope%20is%20outside%20this%20client%27s%20ceiling&iss=https%3A%2F%2Fplytarium.com&state=a%20b%26c%3Dd"
|
||||
if got != want {
|
||||
t.Fatalf("got: %s\nwant: %s", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOrderedRedirectAppendsToExistingQuery(t *testing.T) {
|
||||
got := appendOrderedQuery("https://client.example.test/cb?already=here", [][2]string{
|
||||
{"error", "invalid_request"},
|
||||
{"iss", "https://plytarium.com"},
|
||||
})
|
||||
want := "https://client.example.test/cb?already=here&error=invalid_request&iss=https%3A%2F%2Fplytarium.com"
|
||||
if got != want {
|
||||
t.Fatalf("got: %s\nwant: %s", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeBackendUnavailableIsOpaque500(t *testing.T) {
|
||||
opts := DefaultOptions()
|
||||
opts.Issuer = "https://plytarium.com"
|
||||
srv := NewServer(opts)
|
||||
req := authorizeRequest(map[string]string{"client_id": "anything"})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusInternalServerError {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -57,6 +57,17 @@ type Options struct {
|
||||
// RegisterMaxBodyBytes bounds the RFC 7591 registration request body
|
||||
// before JSON decoding (D-21, T-08-DCR-FLOOD). PHP default: 65536 (64 KiB).
|
||||
RegisterMaxBodyBytes int64
|
||||
|
||||
// Resource is the expected RFC 8707 resource indicator value authorize
|
||||
// checks an optional resource query parameter against (PHP
|
||||
// config('fonoteka.mcp.resource'), D-03). PHP default:
|
||||
// "https://mcp.plytarium.com/mcp".
|
||||
Resource string
|
||||
|
||||
// PendingRequestTTL is how long a pre-consent pending authorization row
|
||||
// created by authorize stays valid (PHP
|
||||
// OAuthCodeManager::PENDING_TTL_SECONDS, D-03). PHP default: 600s.
|
||||
PendingRequestTTL time.Duration
|
||||
}
|
||||
|
||||
// DefaultOptions returns PHP-parity defaults for every metadata option
|
||||
@@ -70,6 +81,8 @@ func DefaultOptions() Options {
|
||||
DCRClientCap: 200,
|
||||
DCRUnconsentedSweepAge: 24 * time.Hour,
|
||||
RegisterMaxBodyBytes: 65536,
|
||||
Resource: "https://mcp.plytarium.com/mcp",
|
||||
PendingRequestTTL: 600 * time.Second,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user