338 lines
11 KiB
Go
338 lines
11 KiB
Go
package oauth_test
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"io"
|
|
"log/slog"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"net/url"
|
|
"strings"
|
|
"testing"
|
|
|
|
"github.com/krow/krow-backend/go-api/internal/auth"
|
|
"github.com/krow/krow-backend/go-api/internal/authctx"
|
|
"github.com/krow/krow-backend/go-api/internal/mcpserver"
|
|
"github.com/krow/krow-backend/go-api/internal/oauth"
|
|
"github.com/krow/krow-backend/go-api/internal/runtime"
|
|
"github.com/krow/krow-backend/go-api/internal/testutil"
|
|
)
|
|
|
|
// The seam, joined.
|
|
//
|
|
// This is the test that matters most in Phase 3, and it is in an EXTERNAL test
|
|
// package (oauth_test) on purpose: it may use only the exported surface, which
|
|
// is exactly what the HTTP layer will use when it wires these two packages
|
|
// together in a later phase. If this compiles and passes, the wiring is a
|
|
// constructor call and nothing else.
|
|
//
|
|
// What it proves end to end, with a real database and no fakes anywhere:
|
|
//
|
|
// OAuth authorization code flow
|
|
// → access token
|
|
// → mcpserver.TokenAuthenticator (the PRODUCTION implementation)
|
|
// → authctx.Identity built from the live user row
|
|
// → tools.Registry.Dispatch
|
|
// → the existing policy table and org pre-filter
|
|
// → real rows from Postgres
|
|
|
|
const (
|
|
itIssuer = "https://api.example.test"
|
|
itResource = "https://api.example.test/mcp"
|
|
itRedirect = "https://claude.example.test/callback"
|
|
itVerifier = "abcdefghijklmnopqrstuvwxyz0123456789ABCDEFG"
|
|
)
|
|
|
|
// itSession is the SessionResolver a signed-in browser satisfies with a cookie.
|
|
// The real implementation lives in the HTTP layer; this stands in for it so the
|
|
// flow can be driven without a browser.
|
|
type itSession struct{ id authctx.Identity }
|
|
|
|
func (s itSession) CurrentUser(*http.Request) (authctx.Identity, bool) { return s.id, true }
|
|
|
|
func sessionFor(userID, orgID string) itSession {
|
|
return itSession{id: authctx.Identity{
|
|
UserID: userID, OrgID: orgID, Role: "admin",
|
|
Email: "a@example.test", Status: "active", AccountType: "employer",
|
|
}}
|
|
}
|
|
|
|
// itRegister performs dynamic client registration over the real handler.
|
|
func itRegister(t *testing.T, as *oauth.Server) string {
|
|
t.Helper()
|
|
body := `{"client_name":"Integration Client","redirect_uris":["` + itRedirect + `"]}`
|
|
req := httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body))
|
|
rec := httptest.NewRecorder()
|
|
as.RegisterHandler().ServeHTTP(rec, req)
|
|
if rec.Code != http.StatusCreated {
|
|
t.Fatalf("register: %d %s", rec.Code, rec.Body.String())
|
|
}
|
|
var out struct {
|
|
ClientID string `json:"client_id"`
|
|
}
|
|
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
|
|
t.Fatalf("register response: %v", err)
|
|
}
|
|
return out.ClientID
|
|
}
|
|
|
|
// itAuthorizeParams is a well-formed authorization request.
|
|
func itAuthorizeParams(clientID string) url.Values {
|
|
return url.Values{
|
|
"client_id": {clientID}, "redirect_uri": {itRedirect}, "response_type": {"code"},
|
|
"state": {"st8"}, "code_challenge": {oauth.ChallengeFor(itVerifier)},
|
|
"code_challenge_method": {"S256"}, "resource": {itResource}, "scope": {oauth.ScopeRead},
|
|
}
|
|
}
|
|
|
|
// itCSRF pulls the consent form's token out of the rendered page.
|
|
func itCSRF(t *testing.T, body string) string {
|
|
t.Helper()
|
|
const marker = `name="csrf" value="`
|
|
i := strings.Index(body, marker)
|
|
if i < 0 {
|
|
t.Fatalf("no csrf field in the consent page")
|
|
}
|
|
rest := body[i+len(marker):]
|
|
return rest[:strings.Index(rest, `"`)]
|
|
}
|
|
|
|
// itDecide posts an approve/deny decision.
|
|
func itDecide(t *testing.T, as *oauth.Server, params url.Values, decision, csrf string) *httptest.ResponseRecorder {
|
|
t.Helper()
|
|
form := url.Values{}
|
|
for k, v := range params {
|
|
form[k] = v
|
|
}
|
|
form.Set("decision", decision)
|
|
form.Set("csrf", csrf)
|
|
req := httptest.NewRequest(http.MethodPost, "/oauth/authorize", strings.NewReader(form.Encode()))
|
|
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
|
rec := httptest.NewRecorder()
|
|
as.AuthorizeHandler().ServeHTTP(rec, req)
|
|
return rec
|
|
}
|
|
|
|
// itAuthorize drives the authorization endpoint through CONSENT and returns
|
|
// the code.
|
|
func itAuthorize(t *testing.T, as *oauth.Server, clientID string) string {
|
|
t.Helper()
|
|
q := itAuthorizeParams(clientID)
|
|
|
|
// The consent page first — a GET no longer issues a code.
|
|
page := httptest.NewRecorder()
|
|
as.AuthorizeHandler().ServeHTTP(page,
|
|
httptest.NewRequest(http.MethodGet, "/oauth/authorize?"+q.Encode(), nil))
|
|
if page.Code != http.StatusOK {
|
|
t.Fatalf("consent page: %d %s", page.Code, page.Body.String())
|
|
}
|
|
|
|
rec := itDecide(t, as, q, "approve", itCSRF(t, page.Body.String()))
|
|
if rec.Code != http.StatusFound {
|
|
t.Fatalf("authorize: %d %s", rec.Code, rec.Body.String())
|
|
}
|
|
loc, err := url.Parse(rec.Header().Get("Location"))
|
|
if err != nil {
|
|
t.Fatalf("Location: %v", err)
|
|
}
|
|
code := loc.Query().Get("code")
|
|
if code == "" {
|
|
t.Fatalf("no code: %s", loc)
|
|
}
|
|
return code
|
|
}
|
|
|
|
// itExchange redeems the code for an access token.
|
|
func itExchange(t *testing.T, as *oauth.Server, clientID, code string) string {
|
|
t.Helper()
|
|
form := url.Values{
|
|
"grant_type": {"authorization_code"}, "code": {code}, "client_id": {clientID},
|
|
"redirect_uri": {itRedirect}, "code_verifier": {itVerifier},
|
|
}
|
|
req := httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(form.Encode()))
|
|
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
|
rec := httptest.NewRecorder()
|
|
as.TokenHandler().ServeHTTP(rec, req)
|
|
if rec.Code != http.StatusOK {
|
|
t.Fatalf("token: %d %s", rec.Code, rec.Body.String())
|
|
}
|
|
var out struct {
|
|
AccessToken string `json:"access_token"`
|
|
}
|
|
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
|
|
t.Fatalf("token response: %v", err)
|
|
}
|
|
return out.AccessToken
|
|
}
|
|
|
|
func TestOAuthTokenReachesMCPToolsAndRealAuthorization(t *testing.T) {
|
|
db := testutil.New(t)
|
|
ctx := context.Background()
|
|
log := slog.New(slog.NewTextHandler(io.Discard, nil))
|
|
|
|
// Two tenants with different volumes, so a leak is visible as a number.
|
|
orgA := mustOrg(t, db, "it-org-a")
|
|
orgB := mustOrg(t, db, "it-org-b")
|
|
userA := mustUser(t, db, orgA, "a@example.test", "admin")
|
|
seedActivity(t, db, orgA, 7, "a@example.test")
|
|
seedActivity(t, db, orgB, 55, "b@example.test")
|
|
|
|
store := oauth.NewStore(db.Pool)
|
|
|
|
// ── Register, authorize, exchange: the real flow, over the real handlers.
|
|
as := oauth.NewServer(
|
|
oauth.Config{Issuer: itIssuer, Resource: itResource},
|
|
store,
|
|
sessionFor(userA, orgA),
|
|
"/login", log,
|
|
)
|
|
|
|
clientID := itRegister(t, as)
|
|
code := itAuthorize(t, as, clientID)
|
|
accessToken := itExchange(t, as, clientID, code)
|
|
|
|
// ── The production authenticator, plugged into the Phase 2 seam.
|
|
authenticator := oauth.NewAuthenticator(store, auth.NewPGUserStore(db.Pool), itResource, log)
|
|
mcp := mcpserver.New(runtime.DefaultTools(db.Pool, nil), authenticator, log)
|
|
|
|
// ── A real MCP tool call, carrying a real OAuth token.
|
|
rec := itCall(t, mcp, "Bearer "+accessToken,
|
|
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
|
{"name":"activity_breakdown","arguments":{}}}`)
|
|
|
|
if rec.Code != http.StatusOK {
|
|
t.Fatalf("MCP call with an OAuth token: %d %s", rec.Code, rec.Body.String())
|
|
}
|
|
total := itTotalEvents(t, rec)
|
|
if total != 7 {
|
|
t.Errorf("totalEvents = %d, want 7 (org A only). Org B has 55; a wrong "+
|
|
"number here means the OAuth identity did not scope the query", total)
|
|
}
|
|
|
|
// ── Without the token, the same call must be refused.
|
|
if rec := itCall(t, mcp, "",
|
|
`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`); rec.Code != http.StatusUnauthorized {
|
|
t.Errorf("unauthenticated MCP call = %d, want 401", rec.Code)
|
|
}
|
|
|
|
// ── Revoking disconnects: the same token must stop working immediately,
|
|
// not at expiry.
|
|
if err := store.RevokeToken(ctx, accessToken, "test_disconnect"); err != nil {
|
|
t.Fatalf("revoke: %v", err)
|
|
}
|
|
if rec := itCall(t, mcp, "Bearer "+accessToken,
|
|
`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`); rec.Code != http.StatusUnauthorized {
|
|
t.Errorf("a revoked token still reached MCP: %d", rec.Code)
|
|
}
|
|
}
|
|
|
|
// A token for another resource must not open the MCP surface, even though this
|
|
// same server minted it. The confused-deputy case, end to end.
|
|
func TestTokenForAnotherResourceCannotReachMCP(t *testing.T) {
|
|
db := testutil.New(t)
|
|
log := slog.New(slog.NewTextHandler(io.Discard, nil))
|
|
|
|
org := mustOrg(t, db, "it-aud-org")
|
|
user := mustUser(t, db, org, "aud@example.test", "admin")
|
|
store := oauth.NewStore(db.Pool)
|
|
|
|
as := oauth.NewServer(
|
|
oauth.Config{Issuer: itIssuer, Resource: itResource},
|
|
store, sessionFor(user, org), "/login", log)
|
|
clientID := itRegister(t, as)
|
|
|
|
pair, err := store.IssuePair(context.Background(), oauth.Token{
|
|
ClientID: clientID, UserID: user, OrgID: org,
|
|
Scopes: []string{oauth.ScopeRead}, Audience: "https://a-different-service.test/mcp",
|
|
}, "")
|
|
if err != nil {
|
|
t.Fatalf("issue: %v", err)
|
|
}
|
|
|
|
mcp := mcpserver.New(
|
|
runtime.DefaultTools(db.Pool, nil),
|
|
oauth.NewAuthenticator(store, auth.NewPGUserStore(db.Pool), itResource, log),
|
|
log)
|
|
|
|
if rec := itCall(t, mcp, "Bearer "+pair.AccessToken,
|
|
`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`); rec.Code != http.StatusUnauthorized {
|
|
t.Errorf("a token for another resource reached MCP: %d", rec.Code)
|
|
}
|
|
}
|
|
|
|
/* ── Helpers ────────────────────────────────────────────────────────────── */
|
|
|
|
func itCall(t *testing.T, s *mcpserver.Server, authHeader, body string) *httptest.ResponseRecorder {
|
|
t.Helper()
|
|
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body))
|
|
req.Header.Set("Content-Type", "application/json")
|
|
if authHeader != "" {
|
|
req.Header.Set("Authorization", authHeader)
|
|
}
|
|
rec := httptest.NewRecorder()
|
|
s.Handler().ServeHTTP(rec, req)
|
|
return rec
|
|
}
|
|
|
|
func itTotalEvents(t *testing.T, rec *httptest.ResponseRecorder) int {
|
|
t.Helper()
|
|
var envelope struct {
|
|
Result struct {
|
|
Content []struct {
|
|
Text string `json:"text"`
|
|
} `json:"content"`
|
|
IsError bool `json:"isError"`
|
|
} `json:"result"`
|
|
}
|
|
if err := json.Unmarshal(rec.Body.Bytes(), &envelope); err != nil {
|
|
t.Fatalf("response: %v", err)
|
|
}
|
|
if envelope.Result.IsError || len(envelope.Result.Content) == 0 {
|
|
t.Fatalf("tool call failed: %s", rec.Body.String())
|
|
}
|
|
var payload struct {
|
|
Data struct {
|
|
TotalEvents int `json:"totalEvents"`
|
|
} `json:"data"`
|
|
}
|
|
if err := json.Unmarshal([]byte(envelope.Result.Content[0].Text), &payload); err != nil {
|
|
t.Fatalf("tool payload: %v", err)
|
|
}
|
|
return payload.Data.TotalEvents
|
|
}
|
|
|
|
func mustOrg(t *testing.T, db *testutil.Harness, slug string) string {
|
|
t.Helper()
|
|
var id string
|
|
if err := db.Pool.QueryRow(context.Background(),
|
|
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
|
|
slug, slug).Scan(&id); err != nil {
|
|
t.Fatalf("org %s: %v", slug, err)
|
|
}
|
|
return id
|
|
}
|
|
|
|
func mustUser(t *testing.T, db *testutil.Harness, orgID, email, role string) string {
|
|
t.Helper()
|
|
var id string
|
|
if err := db.Pool.QueryRow(context.Background(),
|
|
`INSERT INTO users (org_id, email, full_name, role, account_type, status)
|
|
VALUES ($1::uuid, $2, 'Test', $3, 'employer', 'active') RETURNING id::text`,
|
|
orgID, email, role).Scan(&id); err != nil {
|
|
t.Fatalf("user %s: %v", email, err)
|
|
}
|
|
return id
|
|
}
|
|
|
|
func seedActivity(t *testing.T, db *testutil.Harness, orgID string, n int, email string) {
|
|
t.Helper()
|
|
for i := 0; i < n; i++ {
|
|
if _, err := db.Pool.Exec(context.Background(),
|
|
`INSERT INTO user_activity (org_id, event_type, user_email, user_name)
|
|
VALUES ($1::uuid, 'login', $2, 'Someone')`, orgID, email); err != nil {
|
|
t.Fatalf("seed: %v", err)
|
|
}
|
|
}
|
|
}
|