Files
krow_backend/go-api/internal/oauth/mcp_integration_test.go
Aravind f2aa3b3ad8
Some checks failed
CI / fixture (push) Has been cancelled
CI / test (push) Has been cancelled
mcp connection
2026-09-22 10:58:02 +05:30

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)
}
}
}