mcp connection
This commit is contained in:
337
go-api/internal/oauth/mcp_integration_test.go
Normal file
337
go-api/internal/oauth/mcp_integration_test.go
Normal file
@@ -0,0 +1,337 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user