mcp connection
This commit is contained in:
@@ -206,7 +206,7 @@ func (s *Server) handleLogin(w http.ResponseWriter, r *http.Request) {
|
||||
// which is wider, stops one host working through many accounts. They are
|
||||
// separate limiters because they are deliberately different sizes — see the
|
||||
// note on Server.
|
||||
addr := clientAddr(r)
|
||||
addr := s.trust.clientAddr(r)
|
||||
emailKey := strings.ToLower(email)
|
||||
for _, check := range []struct {
|
||||
limiter *attemptLimiter
|
||||
@@ -318,6 +318,85 @@ var publicPaths = map[string]bool{
|
||||
"/health": true,
|
||||
"/api/v1/auth/login": true,
|
||||
"/api/v1/auth/logout": true,
|
||||
|
||||
// ── The OAuth surface for MCP clients ──────────────────────────────────
|
||||
//
|
||||
// Four paths, each public for a specific reason rather than because
|
||||
// "/oauth/*" is convenient. The namespace is deliberately NOT wildcarded:
|
||||
// /oauth/authorize is not here, because it renders a consent screen for a
|
||||
// signed-in person and must keep requiring a session.
|
||||
//
|
||||
// These routes are registered only when OAUTH_ISSUER and MCP_RESOURCE are
|
||||
// configured. Listing them here is harmless otherwise — an unregistered
|
||||
// path still 404s, it simply does so without being asked for a cookie.
|
||||
|
||||
// RFC 9728 and RFC 8414. A client with no token cannot read a document
|
||||
// that requires one, and these are how it discovers where to get a token.
|
||||
// They contain public endpoint URLs and nothing else.
|
||||
"/.well-known/oauth-protected-resource": true,
|
||||
"/.well-known/oauth-authorization-server": true,
|
||||
|
||||
// RFC 7591. A client that has never registered has no credential to
|
||||
// present; that is what dynamic registration is for.
|
||||
"/oauth/register": true,
|
||||
|
||||
// The client authenticates here with an authorization code or a refresh
|
||||
// token in the BODY. This is a back-channel call from the MCP client's own
|
||||
// servers — there is no browser and no cookie to send.
|
||||
"/oauth/token": true,
|
||||
|
||||
// Revocation authenticates by presenting the token being revoked, for the
|
||||
// same back-channel reason.
|
||||
"/oauth/revoke": true,
|
||||
|
||||
// /mcp is listed here, and it is the entry that most deserves explaining,
|
||||
// because "public" is the opposite of what it means for this path.
|
||||
//
|
||||
// The MCP endpoint authenticates its OWN callers, from the Authorization
|
||||
// header, inside mcpserver — every method but the handshake requires a
|
||||
// valid bearer token, and the transport ignores whatever identity this
|
||||
// middleware may have put in the context. So listing it here does not make
|
||||
// it reachable without a credential; it makes THIS middleware step aside
|
||||
// so the one that knows how to answer can.
|
||||
//
|
||||
// It has to step aside. An MCP client discovers how to authenticate by
|
||||
// calling the endpoint with no token and reading the WWW-Authenticate
|
||||
// header of the 401 — RFC 9728, and the first step of the whole flow.
|
||||
// This middleware's 401 carries no such header, so guarding /mcp here
|
||||
// would mean a client received a refusal with nowhere to go and the
|
||||
// connection could never be established. That is not a hypothetical: it is
|
||||
// what TestMCPWithoutBearerReturns401AndDiscoveryPointer caught.
|
||||
//
|
||||
// What stops a cookie authenticating an MCP call is therefore NOT this
|
||||
// allowlist — it is mcpserver taking its identity as a parameter rather
|
||||
// than from the request context. See mcpserver/auth.go, and
|
||||
// TestMCPRejectsACookieSession below.
|
||||
"/mcp": true,
|
||||
|
||||
// /oauth/authorize is here for the same reason as /mcp, and it took a live
|
||||
// client to show why.
|
||||
//
|
||||
// It was withheld on the reasoning that consent needs a signed-in person,
|
||||
// so the route "genuinely wants the cookie". That reasoning was right about
|
||||
// the requirement and wrong about who enforces it. THE HANDLER already
|
||||
// enforces it — authserver.go asks sessions.CurrentUser, refuses to render
|
||||
// consent without an identity, and redirects an anonymous visitor to the
|
||||
// login with the authorization request preserved in returnTo. Guarding the
|
||||
// path HERE meant that handler was never reached, so the redirect it
|
||||
// performs could never run: every signed-out visitor got this middleware's
|
||||
// JSON 401 instead of a login page.
|
||||
//
|
||||
// That is not a cosmetic difference. A first-time connector user is signed
|
||||
// out by definition, so OAuth's browser leg was unreachable for exactly the
|
||||
// people who needed it. Claude Web stopped here — discovery, registration,
|
||||
// then a 401 with nowhere to go. Claude Desktop only got past it because a
|
||||
// session had been established by hand beforehand.
|
||||
//
|
||||
// Listing it grants nothing: no session still means no consent screen and
|
||||
// no authorization code, and the consent POST still requires the
|
||||
// session-bound CSRF token. What changes is only WHICH layer says no, and
|
||||
// therefore whether it can say "sign in here" instead of "no".
|
||||
"/oauth/authorize": true,
|
||||
}
|
||||
|
||||
// authenticate resolves the session cookie into an identity, or refuses.
|
||||
|
||||
217
go-api/internal/httpserver/clientip.go
Normal file
217
go-api/internal/httpserver/clientip.go
Normal file
@@ -0,0 +1,217 @@
|
||||
package httpserver
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Resolving the caller's network address behind a reverse proxy.
|
||||
//
|
||||
// WHAT THIS IS FOR
|
||||
//
|
||||
// Three limits on this API are keyed by the caller's address: failed logins
|
||||
// (auth.go), OAuth client registration, and OAuth authorization before the
|
||||
// caller has signed in. None of them has a better identity available —
|
||||
// registration is anonymous by definition, and a login attempt is anonymous
|
||||
// until the password has been judged.
|
||||
//
|
||||
// Behind a proxy, net/http reports the PROXY's address on every request. Those
|
||||
// three budgets then describe the proxy rather than the caller, which means one
|
||||
// bucket for the whole deployment: one person retrying a connector exhausts
|
||||
// everybody's registration allowance, and twenty failed passwords anywhere lock
|
||||
// out every user's sign-in. That is the fault this file exists to fix.
|
||||
//
|
||||
// WHY IT IS NOT JUST X-Forwarded-For
|
||||
//
|
||||
// The header is written by clients as readily as by proxies. Believing it
|
||||
// unconditionally is worse than the shared bucket rather than better: a caller
|
||||
// who reaches the API directly can put a different value in every request and
|
||||
// get a fresh budget each time, which is not a weakened limit but no limit at
|
||||
// all. The header carries information only about the hop that appended it, so
|
||||
// it is worth exactly as much as the peer that handed it over.
|
||||
//
|
||||
// Hence: believe it only when the immediate peer is a configured proxy, and
|
||||
// walk the chain from the right, where the entries were written by the hops
|
||||
// closest to us, discarding those that are themselves trusted proxies. The
|
||||
// first address that is not one of ours is the nearest thing to the real client
|
||||
// that the topology can actually vouch for. Everything to its left was supplied
|
||||
// by something we do not control and is never read.
|
||||
//
|
||||
// FAILING SAFE
|
||||
//
|
||||
// Every fallback in here returns the PEER address. That is deliberate and it is
|
||||
// the property worth preserving if this code is ever changed: a bad or missing
|
||||
// chain can only ever make a bucket coarser — more callers sharing one budget,
|
||||
// which is the old behaviour — and can never hand a caller a bucket of their
|
||||
// own. Spoofing gains nothing because no path exists from an untrusted input to
|
||||
// a distinct key.
|
||||
|
||||
// proxyTrust turns a request into the address key used for rate limiting.
|
||||
//
|
||||
// A value rather than a package-level variable so that the trusted set is
|
||||
// wired once at construction and cannot be changed by anything holding a
|
||||
// request. An empty proxyTrust is valid and trusts nothing.
|
||||
type proxyTrust struct {
|
||||
// trusted networks, already masked by config parsing.
|
||||
trusted []netip.Prefix
|
||||
}
|
||||
|
||||
// newProxyTrust builds the resolver from configuration.
|
||||
func newProxyTrust(trusted []netip.Prefix) proxyTrust {
|
||||
return proxyTrust{trusted: trusted}
|
||||
}
|
||||
|
||||
// forwardedHeader is the de facto standard, and what Traefik, nginx, Envoy and
|
||||
// the cloud load balancers all append to.
|
||||
//
|
||||
// RFC 7239's `Forwarded:` header is deliberately NOT read. Supporting both
|
||||
// would mean deciding which wins when they disagree, and an attacker choosing
|
||||
// the one this code happens to prefer. One header, one meaning.
|
||||
const forwardedHeader = "X-Forwarded-For"
|
||||
|
||||
// clientAddr returns the rate-limiting key for the caller's address.
|
||||
//
|
||||
// The port is stripped: a browser opens a new source port per connection, so
|
||||
// keying on host:port would give every attempt its own budget and limit nothing
|
||||
// at all. IPv6 is keyed by /64 — see bucketKey.
|
||||
func (t proxyTrust) clientAddr(r *http.Request) string {
|
||||
peer, ok := parseHost(r.RemoteAddr)
|
||||
if !ok {
|
||||
// RemoteAddr is not something this code recognises — a test server with
|
||||
// a synthetic value, or a unix socket. Key by it verbatim, which is
|
||||
// what this function did before proxies were considered at all.
|
||||
return strings.TrimSpace(r.RemoteAddr)
|
||||
}
|
||||
peerKey := bucketKey(peer)
|
||||
|
||||
// Nothing is trusted, so nothing is read. The common case, and the default.
|
||||
if len(t.trusted) == 0 || !t.contains(peer) {
|
||||
return peerKey
|
||||
}
|
||||
|
||||
if client, ok := t.forwardedClient(r); ok {
|
||||
return bucketKey(client)
|
||||
}
|
||||
return peerKey
|
||||
}
|
||||
|
||||
// forwardedClient walks the forwarded chain from the right and returns the
|
||||
// first address that is not one of our own proxies.
|
||||
//
|
||||
// It reports false — meaning "fall back to the peer" — for an absent header, a
|
||||
// chain that is entirely trusted proxies, and a malformed entry. The last of
|
||||
// those is the interesting one: a chain that cannot be parsed cannot be
|
||||
// reasoned about, and the safe reading of "10.0.0.1, ???, 10.0.0.2" is that
|
||||
// everything to the left of the damage is unusable. Skipping the bad entry and
|
||||
// carrying on would let a caller put anything it likes in the header and have
|
||||
// this code step over it to reach the value the caller wanted read.
|
||||
func (t proxyTrust) forwardedClient(r *http.Request) (netip.Addr, bool) {
|
||||
// Values(), not Get(), because a chain may arrive as several headers as
|
||||
// well as one comma-separated list; they are the same list in HTTP's terms
|
||||
// and the rightmost entry of the last header is the most recent hop.
|
||||
var chain []string
|
||||
for _, header := range r.Header.Values(forwardedHeader) {
|
||||
for _, entry := range strings.Split(header, ",") {
|
||||
chain = append(chain, strings.TrimSpace(entry))
|
||||
}
|
||||
}
|
||||
|
||||
for i := len(chain) - 1; i >= 0; i-- {
|
||||
entry := chain[i]
|
||||
if entry == "" {
|
||||
// A stray comma. Treated as damage rather than skipped, for the
|
||||
// reason in the doc comment above.
|
||||
return netip.Addr{}, false
|
||||
}
|
||||
addr, ok := parseForwardedAddr(entry)
|
||||
if !ok {
|
||||
return netip.Addr{}, false
|
||||
}
|
||||
if t.contains(addr) {
|
||||
// One of ours. Keep walking left, towards the client.
|
||||
continue
|
||||
}
|
||||
return addr, true
|
||||
}
|
||||
// Either there was no header, or every hop in it was a trusted proxy and
|
||||
// none of them recorded a client. Neither tells us who called.
|
||||
return netip.Addr{}, false
|
||||
}
|
||||
|
||||
// contains reports whether an address is one of the configured proxies.
|
||||
func (t proxyTrust) contains(addr netip.Addr) bool {
|
||||
addr = addr.Unmap()
|
||||
for _, prefix := range t.trusted {
|
||||
if prefix.Contains(addr) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// bucketKey is the string a rate-limit bucket is keyed by.
|
||||
//
|
||||
// IPv4 keys by the exact address, which is what this service has always done
|
||||
// and what the existing buckets contain.
|
||||
//
|
||||
// IPv6 keys by the /64 PREFIX instead. A single customer is routinely delegated
|
||||
// a whole /64 — often a /56 or shorter — and every address in it is one
|
||||
// machine's to choose. Keying by the full address would hand one caller
|
||||
// 18 quintillion budgets, which is a limit in form only. /64 is the smallest
|
||||
// unit that is reliably one subscriber rather than one interface, so it is the
|
||||
// narrowest honest key.
|
||||
func bucketKey(addr netip.Addr) string {
|
||||
addr = addr.Unmap().WithZone("") // a scope id is local to the host, never a caller identity
|
||||
if addr.Is4() {
|
||||
return addr.String()
|
||||
}
|
||||
prefix, err := addr.Prefix(64)
|
||||
if err != nil {
|
||||
return addr.String()
|
||||
}
|
||||
return prefix.String()
|
||||
}
|
||||
|
||||
// parseHost splits "host:port" and parses the host.
|
||||
//
|
||||
// RemoteAddr always carries a port for TCP, but a test server, a unix socket or
|
||||
// a middleware that rewrote it may not, so a bare address is accepted too.
|
||||
func parseHost(remoteAddr string) (netip.Addr, bool) {
|
||||
raw := strings.TrimSpace(remoteAddr)
|
||||
if raw == "" {
|
||||
return netip.Addr{}, false
|
||||
}
|
||||
if host, _, err := net.SplitHostPort(raw); err == nil {
|
||||
raw = host
|
||||
}
|
||||
addr, err := netip.ParseAddr(strings.Trim(raw, "[]"))
|
||||
if err != nil {
|
||||
return netip.Addr{}, false
|
||||
}
|
||||
return addr, true
|
||||
}
|
||||
|
||||
// parseForwardedAddr parses one entry of an X-Forwarded-For chain.
|
||||
//
|
||||
// Entries are bare addresses by the header's convention, but a port turns up in
|
||||
// practice — some proxies append one, and IPv6 is then bracketed. Both forms
|
||||
// are accepted; anything else is malformed and refused.
|
||||
//
|
||||
// "unknown", the obfuscated identifiers RFC 7239 permits, and empty entries are
|
||||
// all refused rather than skipped: they say the chain is not a list of
|
||||
// addresses, and this code declines to guess which of the remaining entries the
|
||||
// proxy meant.
|
||||
func parseForwardedAddr(entry string) (netip.Addr, bool) {
|
||||
if addr, err := netip.ParseAddr(entry); err == nil {
|
||||
return addr, true
|
||||
}
|
||||
// "[2001:db8::1]:443" or "203.0.113.7:443".
|
||||
if host, _, err := net.SplitHostPort(entry); err == nil {
|
||||
if addr, err := netip.ParseAddr(strings.Trim(host, "[]")); err == nil {
|
||||
return addr, true
|
||||
}
|
||||
}
|
||||
return netip.Addr{}, false
|
||||
}
|
||||
358
go-api/internal/httpserver/clientip_test.go
Normal file
358
go-api/internal/httpserver/clientip_test.go
Normal file
@@ -0,0 +1,358 @@
|
||||
package httpserver
|
||||
|
||||
// Unit tests for client-address resolution.
|
||||
//
|
||||
// An INTERNAL test package (httpserver, not httpserver_test) because proxyTrust
|
||||
// is unexported and deliberately so — the trusted set is wired once at server
|
||||
// construction and there is no reason for anything outside this package to
|
||||
// build one. The rest of the package's tests stay external; this file is the
|
||||
// exception because what is under test is a decision procedure, and testing it
|
||||
// through an HTTP server would obscure which input produced which key.
|
||||
//
|
||||
// THE PROPERTY THESE TESTS EXIST TO DEFEND
|
||||
//
|
||||
// No untrusted input may produce a distinct bucket key. Every failure path must
|
||||
// collapse back to the peer address. A test that asserts a spoofed header is
|
||||
// "ignored" by checking it does not appear is not enough — it must check the
|
||||
// key equals the PEER's key, because two different wrong answers are still two
|
||||
// different buckets, and two buckets is the whole exploit.
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func prefixes(t *testing.T, cidrs ...string) []netip.Prefix {
|
||||
t.Helper()
|
||||
out := make([]netip.Prefix, 0, len(cidrs))
|
||||
for _, c := range cidrs {
|
||||
p, err := netip.ParsePrefix(c)
|
||||
if err != nil {
|
||||
t.Fatalf("bad test CIDR %q: %v", c, err)
|
||||
}
|
||||
out = append(out, p.Masked())
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// request builds a request with a peer address and an optional forwarded chain.
|
||||
// A chain entry of "" means the header is absent.
|
||||
func request(remoteAddr string, forwarded ...string) *http.Request {
|
||||
r := &http.Request{
|
||||
RemoteAddr: remoteAddr,
|
||||
Header: http.Header{},
|
||||
}
|
||||
for _, f := range forwarded {
|
||||
r.Header.Add(forwardedHeader, f)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
/* ── A. A direct client's forwarded header is not read ──────────────────── */
|
||||
|
||||
func TestDirectClientForwardedHeaderIgnored(t *testing.T) {
|
||||
// A proxy IS configured — just not this caller. The caller reaches the API
|
||||
// directly and claims to be somebody else.
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
got := trust.clientAddr(request("203.0.113.9:51000", "198.51.100.7"))
|
||||
|
||||
if want := "203.0.113.9"; got != want {
|
||||
t.Errorf("clientAddr = %q, want %q — a direct caller's X-Forwarded-For was believed", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNoTrustedProxiesConfiguredIgnoresForwarded(t *testing.T) {
|
||||
// The default posture. Nothing is trusted, so nothing is read, and the
|
||||
// behaviour is exactly what it was before this setting existed.
|
||||
trust := newProxyTrust(nil)
|
||||
|
||||
got := trust.clientAddr(request("10.0.0.1:4000", "198.51.100.7"))
|
||||
|
||||
if want := "10.0.0.1"; got != want {
|
||||
t.Errorf("clientAddr = %q, want %q — an unconfigured deployment read a forwarded address", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── B. A trusted proxy's forwarded client is used ──────────────────────── */
|
||||
|
||||
func TestTrustedProxyForwardedClientUsed(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
got := trust.clientAddr(request("10.0.0.1:4000", "203.0.113.9"))
|
||||
|
||||
if want := "203.0.113.9"; got != want {
|
||||
t.Errorf("clientAddr = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// The point of the whole change: two users behind the same proxy get two keys.
|
||||
func TestTrustedProxySeparatesTwoClients(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
a := trust.clientAddr(request("10.0.0.1:4000", "203.0.113.9"))
|
||||
b := trust.clientAddr(request("10.0.0.1:4001", "203.0.113.10"))
|
||||
|
||||
if a == b {
|
||||
t.Fatalf("two clients behind one proxy shared the key %q", a)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── C. Multiple hops, walked right to left ─────────────────────────────── */
|
||||
|
||||
func TestMultipleTrustedHopsSelectsFirstUntrusted(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8", "172.16.0.0/12"))
|
||||
|
||||
// client → edge(172.16.0.5) → internal(10.0.0.1) → us.
|
||||
// Right to left: 10.0.0.1 ours, 172.16.0.5 ours, 203.0.113.9 the client.
|
||||
got := trust.clientAddr(request("10.0.0.1:4000", "203.0.113.9, 172.16.0.5, 10.0.0.1"))
|
||||
|
||||
if want := "203.0.113.9"; got != want {
|
||||
t.Errorf("clientAddr = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// The chain split across several headers is the same chain.
|
||||
func TestChainSplitAcrossHeaders(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
got := trust.clientAddr(request("10.0.0.1:4000", "203.0.113.9", "10.0.0.1"))
|
||||
|
||||
if want := "203.0.113.9"; got != want {
|
||||
t.Errorf("clientAddr = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// Entries to the LEFT of the first untrusted address are never read, whatever
|
||||
// they say. This is what stops a client prepending a forged hop.
|
||||
func TestEntriesLeftOfTheClientAreNotRead(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
// The caller put "1.2.3.4" at the head of the chain hoping to be keyed by
|
||||
// it. The proxy appended the address it actually saw.
|
||||
got := trust.clientAddr(request("10.0.0.1:4000", "1.2.3.4, 203.0.113.9, 10.0.0.1"))
|
||||
|
||||
if want := "203.0.113.9"; got != want {
|
||||
t.Errorf("clientAddr = %q, want %q — a forged leading hop was selected", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── D. Spoofing gains nothing ──────────────────────────────────────────── */
|
||||
|
||||
// The exploit this design exists to prevent: an untrusted caller varying the
|
||||
// header to get a fresh budget per request. Every variation must land on the
|
||||
// SAME key, and that key must be the peer's.
|
||||
func TestUntrustedSpoofingCannotProduceDistinctBuckets(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
spoofs := []string{
|
||||
"1.2.3.4",
|
||||
"5.6.7.8",
|
||||
"10.0.0.1", // claiming to BE the trusted proxy
|
||||
"1.1.1.1, 2.2.2.2, 10.0.0.1", // a whole fabricated chain ending in ours
|
||||
"::1",
|
||||
"2001:db8::1",
|
||||
}
|
||||
|
||||
const peerKey = "203.0.113.9"
|
||||
for _, spoof := range spoofs {
|
||||
got := trust.clientAddr(request("203.0.113.9:51000", spoof))
|
||||
if got != peerKey {
|
||||
t.Errorf("X-Forwarded-For %q produced key %q, want %q — spoofing bought a separate bucket",
|
||||
spoof, got, peerKey)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A trusted proxy that forwards a chain whose leading entries were forged still
|
||||
// yields one key per real client, not one per forgery.
|
||||
func TestSpoofedPrefixBehindTrustedProxyIsStable(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
first := trust.clientAddr(request("10.0.0.1:4000", "9.9.9.9, 203.0.113.9, 10.0.0.1"))
|
||||
second := trust.clientAddr(request("10.0.0.1:4002", "8.8.8.8, 203.0.113.9, 10.0.0.1"))
|
||||
|
||||
if first != second {
|
||||
t.Errorf("one client produced two keys (%q, %q) by varying a forged hop", first, second)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── E. Malformed input falls back, and never panics ────────────────────── */
|
||||
|
||||
func TestMalformedForwardedEntriesFallBackToPeer(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
cases := map[string]string{
|
||||
"not an address": "banana",
|
||||
"unknown": "unknown",
|
||||
"obfuscated (7239)": "_hidden",
|
||||
"empty entry": "203.0.113.9, , 10.0.0.1",
|
||||
"trailing comma": "203.0.113.9,",
|
||||
"damage before ours": "203.0.113.9, banana, 10.0.0.1",
|
||||
"whitespace only": " ",
|
||||
"port but no host": ":443",
|
||||
"cidr not address": "203.0.113.0/24",
|
||||
}
|
||||
|
||||
const peerKey = "10.0.0.1"
|
||||
for name, header := range cases {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
got := trust.clientAddr(request("10.0.0.1:4000", header))
|
||||
if got != peerKey {
|
||||
t.Errorf("clientAddr = %q, want the peer %q", got, peerKey)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// An address WITH a port is not malformed — some proxies append one.
|
||||
func TestForwardedEntryWithPortIsAccepted(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
if got, want := trust.clientAddr(request("10.0.0.1:4000", "203.0.113.9:51000")), "203.0.113.9"; got != want {
|
||||
t.Errorf("IPv4 with port: clientAddr = %q, want %q", got, want)
|
||||
}
|
||||
if got, want := trust.clientAddr(request("10.0.0.1:4000", "[2001:db8::1]:443")), "2001:db8::/64"; got != want {
|
||||
t.Errorf("IPv6 with port: clientAddr = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMalformedRemoteAddrDoesNotPanic(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
for _, remote := range []string{"", " ", "pipe", "not:an:addr", "@"} {
|
||||
got := trust.clientAddr(request(remote, "203.0.113.9"))
|
||||
// Whatever it returns, it must not be the forwarded address: an
|
||||
// unparseable peer is not a trusted one.
|
||||
if got == "203.0.113.9" {
|
||||
t.Errorf("RemoteAddr %q was treated as a trusted peer", remote)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── F. IPv6 is keyed by /64 ────────────────────────────────────────────── */
|
||||
|
||||
func TestIPv6SameSlash64SharesABucket(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
// Same /64, different hosts within it — one subscriber, one budget.
|
||||
a := trust.clientAddr(request("10.0.0.1:4000", "2001:db8:abcd:1234::1"))
|
||||
b := trust.clientAddr(request("10.0.0.1:4000", "2001:db8:abcd:1234:ffff:ffff:ffff:ffff"))
|
||||
|
||||
if a != b {
|
||||
t.Errorf("two addresses in one /64 produced %q and %q; a caller could mint budgets at will", a, b)
|
||||
}
|
||||
if want := "2001:db8:abcd:1234::/64"; a != want {
|
||||
t.Errorf("key = %q, want %q", a, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIPv6DifferentSlash64DoesNotShareABucket(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
a := trust.clientAddr(request("10.0.0.1:4000", "2001:db8:abcd:1234::1"))
|
||||
b := trust.clientAddr(request("10.0.0.1:4000", "2001:db8:abcd:9999::1"))
|
||||
|
||||
if a == b {
|
||||
t.Errorf("two different /64s shared the key %q", a)
|
||||
}
|
||||
}
|
||||
|
||||
// An IPv4 peer reported in IPv4-mapped form is the same caller as the plain
|
||||
// form, and must not become a second bucket.
|
||||
func TestIPv4MappedIPv6NormalisesToIPv4(t *testing.T) {
|
||||
trust := newProxyTrust(nil)
|
||||
|
||||
plain := trust.clientAddr(request("203.0.113.9:51000"))
|
||||
mapped := trust.clientAddr(request("[::ffff:203.0.113.9]:51000"))
|
||||
|
||||
if plain != mapped {
|
||||
t.Errorf("plain %q and mapped %q are the same host but keyed differently", plain, mapped)
|
||||
}
|
||||
if want := "203.0.113.9"; plain != want {
|
||||
t.Errorf("key = %q, want %q", plain, want)
|
||||
}
|
||||
}
|
||||
|
||||
// A trusted IPv6 proxy works the same way as a trusted IPv4 one.
|
||||
func TestTrustedIPv6Proxy(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "fd00::/8"))
|
||||
|
||||
got := trust.clientAddr(request("[fd00::1]:4000", "2001:db8:abcd:1234::5"))
|
||||
|
||||
if want := "2001:db8:abcd:1234::/64"; got != want {
|
||||
t.Errorf("clientAddr = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// A scope id is local to this host and says nothing about who called.
|
||||
func TestIPv6ZoneIsNotPartOfTheKey(t *testing.T) {
|
||||
trust := newProxyTrust(nil)
|
||||
|
||||
withZone := trust.clientAddr(request("[fe80::1%eth0]:4000"))
|
||||
without := trust.clientAddr(request("[fe80::1]:4000"))
|
||||
|
||||
if withZone != without {
|
||||
t.Errorf("zone changed the key: %q vs %q", withZone, without)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── G. No header at all ────────────────────────────────────────────────── */
|
||||
|
||||
func TestMissingForwardedHeaderFallsBackToPeer(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
if got, want := trust.clientAddr(request("10.0.0.1:4000")), "10.0.0.1"; got != want {
|
||||
t.Errorf("clientAddr = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// A chain consisting only of our own proxies names no client.
|
||||
func TestChainOfOnlyTrustedProxiesFallsBackToPeer(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "10.0.0.0/8"))
|
||||
|
||||
if got, want := trust.clientAddr(request("10.0.0.1:4000", "10.0.0.2, 10.0.0.1")), "10.0.0.1"; got != want {
|
||||
t.Errorf("clientAddr = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── H. The port is not part of the key ─────────────────────────────────── */
|
||||
|
||||
// Pre-existing behaviour, asserted here because it is the reason this function
|
||||
// strips the port at all: a browser opens a new source port per connection.
|
||||
func TestSourcePortIsNotPartOfTheKey(t *testing.T) {
|
||||
trust := newProxyTrust(nil)
|
||||
|
||||
a := trust.clientAddr(request("203.0.113.9:51000"))
|
||||
b := trust.clientAddr(request("203.0.113.9:51001"))
|
||||
|
||||
if a != b {
|
||||
t.Errorf("source port changed the key: %q vs %q", a, b)
|
||||
}
|
||||
}
|
||||
|
||||
// A bare address with no port — a test server, or a rewritten RemoteAddr.
|
||||
func TestRemoteAddrWithoutAPortIsAccepted(t *testing.T) {
|
||||
trust := newProxyTrust(nil)
|
||||
|
||||
if got, want := trust.clientAddr(request("203.0.113.9")), "203.0.113.9"; got != want {
|
||||
t.Errorf("clientAddr = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Trust-set edge cases ───────────────────────────────────────────────── */
|
||||
|
||||
// A single-host trusted proxy, which is what a bare address in configuration
|
||||
// becomes.
|
||||
func TestSingleHostTrustedProxy(t *testing.T) {
|
||||
trust := newProxyTrust(prefixes(t, "172.17.0.1/32"))
|
||||
|
||||
if got, want := trust.clientAddr(request("172.17.0.1:4000", "203.0.113.9")), "203.0.113.9"; got != want {
|
||||
t.Errorf("trusted host: clientAddr = %q, want %q", got, want)
|
||||
}
|
||||
// One address along is NOT trusted.
|
||||
if got, want := trust.clientAddr(request("172.17.0.2:4000", "203.0.113.9")), "172.17.0.2"; got != want {
|
||||
t.Errorf("neighbouring host: clientAddr = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
193
go-api/internal/httpserver/maintenance.go
Normal file
193
go-api/internal/httpserver/maintenance.go
Normal file
@@ -0,0 +1,193 @@
|
||||
package httpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/oauth"
|
||||
"github.com/krow/krow-backend/go-api/internal/ratelimit"
|
||||
)
|
||||
|
||||
// Scheduled maintenance for the OAuth and rate-limit tables.
|
||||
//
|
||||
// WHY THIS SHAPE AND NOT A NEW ONE
|
||||
//
|
||||
// The process already has a scheduled maintenance mechanism: sweepSessions in
|
||||
// cmd/api/main.go, a ticker goroutine whose context is the server's, which runs
|
||||
// once at startup and then on an interval, logs a failure and retries at the
|
||||
// next tick. It is bounded, cancellable, non-blocking and failure-isolated, and
|
||||
// it has been in production.
|
||||
//
|
||||
// So this is the same thing for two more tables rather than a second kind of
|
||||
// thing. No new process, no cron dependency, no leader election, no library.
|
||||
// The one addition is that both sweeps live behind a single type, so
|
||||
// cmd/api/main.go gains one line rather than two more goroutines.
|
||||
//
|
||||
// MULTI-INSTANCE SAFETY COMES FROM THE STATEMENTS, NOT FROM COORDINATION
|
||||
//
|
||||
// Every instance runs this, on its own schedule, with no lock between them —
|
||||
// deliberately. A lease or an advisory lock would be state to hold, to expire
|
||||
// and to recover when the holder dies mid-sweep, in exchange for avoiding work
|
||||
// that is already harmless: each sweep is a bounded DELETE whose predicate no
|
||||
// longer matches once a row is gone. Two instances sweeping at the same moment
|
||||
// delete disjoint sets and neither errors. A row deleted twice is not an error;
|
||||
// it is a row that was already deleted.
|
||||
//
|
||||
// That is the same property Phase 5's concurrent-cleanup test asserts directly:
|
||||
// four workers, six dead tokens, exactly six removed between them.
|
||||
|
||||
// maintenanceInterval is how often the sweep runs.
|
||||
//
|
||||
// Hourly. The grace period before anything is deleted is also an hour, so a
|
||||
// row becomes eligible and is collected within roughly two — soon enough that
|
||||
// nothing accumulates, and far enough apart that a DELETE never lands on a hot
|
||||
// path. Shorter would buy nothing: nothing here is a correctness deadline.
|
||||
//
|
||||
// Deliberately NOT sweepInterval's fifteen minutes. Sessions churn with every
|
||||
// sign-in; authorization codes live sixty seconds and tokens fifteen minutes,
|
||||
// so an hour still collects them promptly while running a quarter as often.
|
||||
const maintenanceInterval = time.Hour
|
||||
|
||||
// maintenanceTimeout bounds one pass.
|
||||
//
|
||||
// Generous for three bounded deletes and short enough that a wedged statement
|
||||
// cannot hold this goroutine past shutdown. Matches sweepSessions' own bound in
|
||||
// spirit; longer only because there are more statements.
|
||||
const maintenanceTimeout = 60 * time.Second
|
||||
|
||||
// Maintenance sweeps the OAuth and rate-limit tables.
|
||||
//
|
||||
// Nil when the deployment does not serve MCP, which is why Server.Maintenance
|
||||
// returns a pointer and the caller checks it — the same way routeOAuth simply
|
||||
// registers nothing.
|
||||
type Maintenance struct {
|
||||
store *oauth.Store
|
||||
limiter *ratelimit.Limiter
|
||||
log *slog.Logger
|
||||
}
|
||||
|
||||
// Maintenance exposes the sweeper, or nil when there is nothing to sweep.
|
||||
//
|
||||
// Mirrors Server.Sessions(), which exists for exactly this reason: the process
|
||||
// owns the schedule, the server owns the things being swept.
|
||||
func (s *Server) Maintenance() *Maintenance {
|
||||
if !s.cfg.OAuth.Enabled() {
|
||||
return nil
|
||||
}
|
||||
return &Maintenance{
|
||||
store: oauth.NewStore(s.db.Pool),
|
||||
limiter: s.limiter,
|
||||
log: s.log,
|
||||
}
|
||||
}
|
||||
|
||||
// MaintenanceResult is what one pass removed.
|
||||
type MaintenanceResult struct {
|
||||
Grants int64
|
||||
AccessTokens int64
|
||||
RefreshTokens int64
|
||||
RateLimits int64
|
||||
}
|
||||
|
||||
// Total is the row count removed, for the log line.
|
||||
func (r MaintenanceResult) Total() int64 {
|
||||
return r.Grants + r.AccessTokens + r.RefreshTokens + r.RateLimits
|
||||
}
|
||||
|
||||
// Sweep runs one maintenance pass.
|
||||
//
|
||||
// The two halves are independent on purpose: a failure sweeping OAuth rows must
|
||||
// not prevent the rate-limit sweep, because the second is the one that would
|
||||
// otherwise grow without bound. The first error is returned, after both have
|
||||
// been attempted.
|
||||
func (m *Maintenance) Sweep(ctx context.Context) (MaintenanceResult, error) {
|
||||
var out MaintenanceResult
|
||||
var firstErr error
|
||||
|
||||
// OAuth: codes, access tokens, and refresh tokens past their retention.
|
||||
// The grace period and the reuse-detection retention are enforced inside
|
||||
// Store.Cleanup — this schedules it, it does not reimplement it.
|
||||
cleaned, err := m.store.Cleanup(ctx)
|
||||
if err != nil {
|
||||
firstErr = err
|
||||
} else {
|
||||
out.Grants = cleaned.Grants
|
||||
out.AccessTokens = cleaned.AccessTokens
|
||||
out.RefreshTokens = cleaned.RefreshTokens
|
||||
}
|
||||
|
||||
if m.limiter != nil {
|
||||
swept, err := m.limiter.Sweep(ctx, 0) // 0 = the package's own batch size
|
||||
if err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
out.RateLimits = swept
|
||||
}
|
||||
|
||||
return out, firstErr
|
||||
}
|
||||
|
||||
// SweepMaintenance runs the sweep until the context is cancelled.
|
||||
//
|
||||
// Deliberately identical in shape to sweepSessions: one pass immediately so a
|
||||
// process that has been down does not carry a backlog for a further hour, then
|
||||
// on the ticker. A failed pass is logged and retried at the next tick — the
|
||||
// tables being briefly larger than they should be is not worth stopping the API
|
||||
// for, and it is certainly not worth a panic in a goroutine nobody is watching.
|
||||
//
|
||||
// Exported because cmd/api owns the process's goroutines and this package owns
|
||||
// what they do.
|
||||
func SweepMaintenance(ctx context.Context, m *Maintenance, log *slog.Logger) {
|
||||
if m == nil {
|
||||
// No OAuth surface, nothing to sweep. Returning rather than ticking
|
||||
// uselessly for the life of the process.
|
||||
return
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(maintenanceInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
pass := func() {
|
||||
// A deadline of its own, so a slow DELETE cannot leave this goroutine
|
||||
// blocked past shutdown.
|
||||
sweepCtx, cancel := context.WithTimeout(ctx, maintenanceTimeout)
|
||||
defer cancel()
|
||||
|
||||
// A panic in a background goroutine takes the process with it, and
|
||||
// this one runs unattended for the life of the deployment. Recovering
|
||||
// turns a bug here into a logged failure and a retry at the next tick.
|
||||
defer func() {
|
||||
if p := recover(); p != nil {
|
||||
log.Error("maintenance sweep panicked", "panic", p)
|
||||
}
|
||||
}()
|
||||
|
||||
result, err := m.Sweep(sweepCtx)
|
||||
switch {
|
||||
case err != nil && ctx.Err() != nil:
|
||||
// Shutting down; the cancellation is expected, not a failure.
|
||||
case err != nil:
|
||||
log.Warn("maintenance sweep failed", "error", err,
|
||||
"grants", result.Grants, "access_tokens", result.AccessTokens,
|
||||
"refresh_tokens", result.RefreshTokens, "rate_limits", result.RateLimits)
|
||||
case result.Total() > 0:
|
||||
log.Info("maintenance sweep",
|
||||
"grants", result.Grants, "access_tokens", result.AccessTokens,
|
||||
"refresh_tokens", result.RefreshTokens, "rate_limits", result.RateLimits)
|
||||
default:
|
||||
log.Debug("maintenance sweep found nothing to delete")
|
||||
}
|
||||
}
|
||||
|
||||
pass()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
log.Debug("maintenance sweeper stopped")
|
||||
return
|
||||
case <-ticker.C:
|
||||
pass()
|
||||
}
|
||||
}
|
||||
}
|
||||
275
go-api/internal/httpserver/maintenance_test.go
Normal file
275
go-api/internal/httpserver/maintenance_test.go
Normal file
@@ -0,0 +1,275 @@
|
||||
package httpserver_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/httpserver"
|
||||
)
|
||||
|
||||
// Scheduler lifecycle.
|
||||
//
|
||||
// What is under test is the GOROUTINE, not the deletes — those are covered in
|
||||
// internal/oauth and internal/ratelimit against real data. Here the questions
|
||||
// are: does it start, does it do a pass, does it stop when told, does a failure
|
||||
// take the process with it, and is running it twice safe.
|
||||
|
||||
/* ── Lifecycle ──────────────────────────────────────────────────────────── */
|
||||
|
||||
// It runs one pass IMMEDIATELY, before the first tick. A process that has been
|
||||
// down should not carry a backlog for a further hour.
|
||||
func TestMaintenanceRunsOnceImmediately(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
m := a.srv.Maintenance()
|
||||
if m == nil {
|
||||
t.Fatal("a configured deployment returned no Maintenance")
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
httpserver.SweepMaintenance(ctx, m, slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// The immediate pass is the only one that will happen inside the test's
|
||||
// lifetime — the ticker is an hour. Give it a moment, then stop.
|
||||
time.Sleep(200 * time.Millisecond)
|
||||
cancel()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("the sweeper did not stop within 5s of cancellation")
|
||||
}
|
||||
}
|
||||
|
||||
// Cancellation must return promptly, or a shutdown hangs on a goroutine nobody
|
||||
// is waiting for.
|
||||
func TestMaintenanceStopsOnCancellation(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
httpserver.SweepMaintenance(ctx, a.srv.Maintenance(),
|
||||
slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
close(done)
|
||||
}()
|
||||
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
start := time.Now()
|
||||
cancel()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
if elapsed := time.Since(start); elapsed > 2*time.Second {
|
||||
t.Errorf("stopping took %v; shutdown would block on it", elapsed)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("the sweeper ignored cancellation")
|
||||
}
|
||||
}
|
||||
|
||||
// An already-cancelled context must not run a pass and must return at once.
|
||||
func TestMaintenanceWithAnAlreadyCancelledContextReturns(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
httpserver.SweepMaintenance(ctx, a.srv.Maintenance(),
|
||||
slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
close(done)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("the sweeper did not return on an already-cancelled context")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── It does the work ───────────────────────────────────────────────────── */
|
||||
|
||||
// One pass removes dead rows and leaves live ones. The detailed retention rules
|
||||
// are tested in internal/oauth; this asserts the scheduler is wired to them.
|
||||
func TestMaintenanceSweepRemovesDeadRows(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// A grant that is already past its expiry and its grace.
|
||||
if _, err := a.h.Pool.Exec(ctx,
|
||||
`INSERT INTO oauth_clients (client_id, client_name, redirect_uris)
|
||||
VALUES ('sweep-client', 'Sweep', ARRAY['https://a.test/cb'])`); err != nil {
|
||||
t.Fatalf("client: %v", err)
|
||||
}
|
||||
userID, _ := seededUser(t, a.h.Pool)
|
||||
var orgID string
|
||||
if err := a.h.Pool.QueryRow(ctx,
|
||||
`SELECT org_id::text FROM users WHERE id = $1::uuid`, userID).Scan(&orgID); err != nil {
|
||||
t.Fatalf("org: %v", err)
|
||||
}
|
||||
|
||||
if _, err := a.h.Pool.Exec(ctx,
|
||||
`INSERT INTO oauth_grants
|
||||
(code_hash, client_id, user_id, org_id, redirect_uri, scopes, resource,
|
||||
code_challenge, code_challenge_method, created_date, expires_at)
|
||||
VALUES (repeat('a', 64), 'sweep-client', $1::uuid, $2::uuid, 'https://a.test/cb',
|
||||
ARRAY['krow.read'], $3, repeat('B', 43), 'S256',
|
||||
now() - interval '3 hours', now() - interval '3 hours' + interval '1 minute')`,
|
||||
userID, orgID, testMCPResource); err != nil {
|
||||
t.Fatalf("grant: %v", err)
|
||||
}
|
||||
|
||||
// An expired rate-limit bucket.
|
||||
if _, err := a.h.Pool.Exec(ctx,
|
||||
`INSERT INTO rate_limits (bucket, window_start, count, expires_at)
|
||||
VALUES ('test:old', now() - interval '3 hours', 5, now() - interval '2 hours')`); err != nil {
|
||||
t.Fatalf("bucket: %v", err)
|
||||
}
|
||||
// And a live one, which must survive.
|
||||
if _, err := a.h.Pool.Exec(ctx,
|
||||
`INSERT INTO rate_limits (bucket, window_start, count, expires_at)
|
||||
VALUES ('test:live', now(), 1, now() + interval '1 hour')`); err != nil {
|
||||
t.Fatalf("bucket: %v", err)
|
||||
}
|
||||
|
||||
result, err := a.srv.Maintenance().Sweep(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Sweep: %v", err)
|
||||
}
|
||||
|
||||
if result.Grants != 1 {
|
||||
t.Errorf("removed %d grants, want 1", result.Grants)
|
||||
}
|
||||
if result.RateLimits != 1 {
|
||||
t.Errorf("removed %d rate-limit rows, want 1", result.RateLimits)
|
||||
}
|
||||
if result.Total() != 2 {
|
||||
t.Errorf("Total() = %d, want 2", result.Total())
|
||||
}
|
||||
|
||||
var live int
|
||||
if err := a.h.Pool.QueryRow(ctx,
|
||||
`SELECT count(*) FROM rate_limits WHERE bucket = 'test:live'`).Scan(&live); err != nil {
|
||||
t.Fatalf("count: %v", err)
|
||||
}
|
||||
if live != 1 {
|
||||
t.Error("the live rate-limit window was swept")
|
||||
}
|
||||
}
|
||||
|
||||
// Running it repeatedly must be safe and must converge to removing nothing.
|
||||
func TestRepeatedMaintenanceIsSafe(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
ctx := context.Background()
|
||||
m := a.srv.Maintenance()
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
result, err := m.Sweep(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("pass %d: %v", i+1, err)
|
||||
}
|
||||
if i > 0 && result.Total() != 0 {
|
||||
t.Errorf("pass %d removed %d rows; a repeat pass should find nothing", i+1, result.Total())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Two instances sweep concurrently with no coordination. Neither may error.
|
||||
// Run with -race.
|
||||
func TestConcurrentMaintenanceIsSafe(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
ctx := context.Background()
|
||||
|
||||
const instances = 4
|
||||
var wg sync.WaitGroup
|
||||
errs := make(chan error, instances)
|
||||
|
||||
for i := 0; i < instances; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
if _, err := a.srv.Maintenance().Sweep(ctx); err != nil {
|
||||
errs <- err
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
close(errs)
|
||||
|
||||
for err := range errs {
|
||||
t.Errorf("concurrent sweep errored: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Failure isolation ──────────────────────────────────────────────────── */
|
||||
|
||||
// A failing sweep must be logged and survived, not fatal. The database is
|
||||
// closed underneath the sweeper, which is the closest thing to a real outage a
|
||||
// test can arrange.
|
||||
func TestMaintenanceSurvivesADatabaseFailure(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
m := a.srv.Maintenance()
|
||||
|
||||
var logged strings.Builder
|
||||
log := slog.New(slog.NewTextHandler(&logged, &slog.HandlerOptions{Level: slog.LevelDebug}))
|
||||
|
||||
// A cancelled context makes every statement fail immediately.
|
||||
dead, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
if _, err := m.Sweep(dead); err == nil {
|
||||
t.Log("note: the sweep reported no error on a cancelled context")
|
||||
}
|
||||
|
||||
// The goroutine wrapper must not panic or exit the process on that.
|
||||
ctx, stop := context.WithCancel(context.Background())
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
httpserver.SweepMaintenance(ctx, m, log)
|
||||
close(done)
|
||||
}()
|
||||
time.Sleep(150 * time.Millisecond)
|
||||
stop()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("the sweeper did not stop")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── It is absent when the surface is ───────────────────────────────────── */
|
||||
|
||||
// A deployment without OAuth has nothing to sweep, and must not start a ticker
|
||||
// that runs for the life of the process doing nothing.
|
||||
func TestMaintenanceIsNilWhenTheSurfaceIsDisabled(t *testing.T) {
|
||||
a := newAPI(t) // the standard fixture: no OAuth configuration
|
||||
|
||||
if m := a.srv.Maintenance(); m != nil {
|
||||
t.Error("an unconfigured deployment returned a Maintenance sweeper")
|
||||
}
|
||||
|
||||
// And the runner must return immediately rather than tick forever.
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
httpserver.SweepMaintenance(context.Background(), a.srv.Maintenance(),
|
||||
slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
close(done)
|
||||
}()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("the sweeper ticked despite having nothing to sweep")
|
||||
}
|
||||
}
|
||||
195
go-api/internal/httpserver/mcp.go
Normal file
195
go-api/internal/httpserver/mcp.go
Normal file
@@ -0,0 +1,195 @@
|
||||
package httpserver
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"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/ratelimit"
|
||||
"github.com/krow/krow-backend/go-api/internal/runtime"
|
||||
)
|
||||
|
||||
// Mounting the MCP surface and the OAuth authorization server behind it.
|
||||
//
|
||||
// This file is the seam between the existing HTTP server and two packages that
|
||||
// know nothing about it. It is deliberately thin: no validation, no policy and
|
||||
// no business logic live here, because every one of those already lives in the
|
||||
// package being mounted. What this file decides is only WHERE things are served
|
||||
// and WHAT AUTHENTICATES them, and those two decisions are the ones that have
|
||||
// to be right.
|
||||
//
|
||||
// OFF UNLESS CONFIGURED. Without OAUTH_ISSUER and MCP_RESOURCE, none of these
|
||||
// routes are registered at all. That follows routeRuns' precedent exactly: a
|
||||
// deployment that does not serve agents answers 404 rather than registering
|
||||
// routes that fail, and the same is true of one that does not serve MCP. An
|
||||
// existing deployment that upgrades to this build gains nothing it did not ask
|
||||
// for.
|
||||
|
||||
// routeOAuth registers the authorization server and its discovery documents.
|
||||
//
|
||||
// WHICH OF THESE ARE PUBLIC, AND WHY — this is the part worth reading twice.
|
||||
// Four paths bypass the cookie middleware, and each has a specific reason:
|
||||
//
|
||||
// /.well-known/oauth-protected-resource RFC 9728. A client that has no
|
||||
// /.well-known/oauth-authorization-server RFC 8414. token cannot read a
|
||||
// document that requires one, and
|
||||
// these are how it learns where to
|
||||
// get a token. They contain only
|
||||
// public endpoint URLs.
|
||||
//
|
||||
// /oauth/register RFC 7591. A client that has never registered has no
|
||||
// credential to present — that is the entire point of
|
||||
// dynamic registration.
|
||||
//
|
||||
// /oauth/token The client authenticates with an authorization code or a
|
||||
// refresh token IN THE BODY. A cookie would be meaningless:
|
||||
// this is a back-channel call from Claude's servers, where
|
||||
// no browser and no cookie exist.
|
||||
//
|
||||
// /oauth/authorize is deliberately NOT public. It runs in a browser, as a
|
||||
// person, and it requires the existing KROW session — that is how the consent
|
||||
// screen knows whose organisation is being granted. An unauthenticated visitor
|
||||
// is redirected to the existing login and comes back.
|
||||
//
|
||||
// /mcp is deliberately NOT public either, and also does not use the cookie. See
|
||||
// routeMCP.
|
||||
func (s *Server) routeOAuth(mux *http.ServeMux) int {
|
||||
if !s.cfg.OAuth.Enabled() {
|
||||
return 0
|
||||
}
|
||||
|
||||
cfg := oauth.Config{
|
||||
Issuer: s.cfg.OAuth.Issuer,
|
||||
Resource: s.cfg.OAuth.Resource,
|
||||
}
|
||||
store := oauth.NewStore(s.db.Pool)
|
||||
as := oauth.NewServer(cfg, store, sessionResolver{s}, s.cfg.OAuth.LoginPath, s.log)
|
||||
|
||||
mux.Handle("GET /.well-known/oauth-protected-resource", cfg.ProtectedResourceHandler())
|
||||
mux.Handle("GET /.well-known/oauth-authorization-server", cfg.AuthorizationServerHandler())
|
||||
// Registration is the only endpoint that writes for a caller with no
|
||||
// credential at all, so it carries the tightest limit on the surface.
|
||||
mux.Handle("POST /oauth/register",
|
||||
s.limited(ratelimit.OAuthRegister, s.byClientAddr, as.RegisterHandler()))
|
||||
// GET renders consent; POST carries the decision. One handler, because the
|
||||
// POST re-validates every parameter the GET validated rather than trusting
|
||||
// the form it rendered.
|
||||
mux.Handle("GET /oauth/authorize",
|
||||
s.limited(ratelimit.OAuthAuthorize, s.byAddrAndUser, as.AuthorizeHandler()))
|
||||
mux.Handle("POST /oauth/authorize",
|
||||
s.limited(ratelimit.OAuthAuthorize, s.byAddrAndUser, as.AuthorizeHandler()))
|
||||
|
||||
// The token endpoint carries two limits on two different subjects, because
|
||||
// its two grant types are abused differently: a code exchange is bounded
|
||||
// per client, and a refresh is bounded per token so a loop on one
|
||||
// connection cannot spend another's budget. Which applies is decided per
|
||||
// request by the grant_type, inside tokenLimited.
|
||||
mux.Handle("POST /oauth/token", s.tokenLimited(as.TokenHandler()))
|
||||
|
||||
// Revocation is deliberately unlimited — see ratelimit/rules.go. It is the
|
||||
// emergency brake, and an attacker gains nothing by pulling it.
|
||||
mux.Handle("POST /oauth/revoke", as.RevokeHandler())
|
||||
|
||||
return 7
|
||||
}
|
||||
|
||||
// routeMCP registers the MCP endpoint.
|
||||
//
|
||||
// AUTHENTICATION HERE IS THE BEARER PATH AND ONLY THE BEARER PATH.
|
||||
//
|
||||
// The handler authenticates its own callers from the Authorization header and
|
||||
// ignores whatever the cookie middleware put in the context. That is a property
|
||||
// of mcpserver, not of this file — see its auth.go.
|
||||
//
|
||||
// /mcp IS on the publicPaths allowlist, and that is deliberate rather than an
|
||||
// oversight. The cookie middleware has to step aside here: an MCP client
|
||||
// discovers how to authenticate by calling this endpoint without a token and
|
||||
// reading the WWW-Authenticate header of the 401, and the middleware's own 401
|
||||
// carries no such header. Guarding the path here would refuse the client with
|
||||
// nowhere to go, and the connection could never be made at all.
|
||||
//
|
||||
// The credential requirement is not weakened by that, because it was never
|
||||
// this middleware enforcing it: mcpserver refuses every method but the
|
||||
// handshake without a bearer token, and it takes its identity as a parameter
|
||||
// rather than from the request context, so a cookie cannot supply one.
|
||||
//
|
||||
// There is no second authorization layer. A tool call goes straight into the
|
||||
// registry the agent runtime already uses, under the policy table it already
|
||||
// consults.
|
||||
func (s *Server) routeMCP(mux *http.ServeMux) int {
|
||||
if !s.cfg.OAuth.Enabled() {
|
||||
return 0
|
||||
}
|
||||
|
||||
// The SAME registry the runtime builds. Not a copy, not a second
|
||||
// construction: a tool added once is available to Owliver and to MCP
|
||||
// together, and neither can drift from the other.
|
||||
registry := runtime.DefaultTools(
|
||||
s.db.Pool,
|
||||
nil, // knowledge_search is not exposed over MCP — see mcpserver/tools.go
|
||||
)
|
||||
|
||||
authenticator := oauth.NewAuthenticator(
|
||||
oauth.NewStore(s.db.Pool),
|
||||
s.users,
|
||||
// The audience an access token must carry. From configuration, never
|
||||
// from a request: a resource value supplied by a caller would let the
|
||||
// caller choose their own audience.
|
||||
s.cfg.OAuth.Resource,
|
||||
s.log,
|
||||
)
|
||||
|
||||
server := mcpserver.New(registry, authenticator, s.log).
|
||||
WithResourceMetadataURL(s.cfg.OAuth.Issuer + "/.well-known/oauth-protected-resource").
|
||||
// The per-organisation ceiling is installed INSIDE the MCP server
|
||||
// rather than as middleware, because the organisation is only known
|
||||
// after the token has been resolved. See orgLimiter in mcplimit.go.
|
||||
WithOrgLimiter(orgLimiter{s})
|
||||
|
||||
mux.Handle("POST /mcp", s.mcpLimited(server.Handler()))
|
||||
// GET is what the Streamable HTTP binding uses for a server-initiated
|
||||
// stream, which this server does not open. Registered so the answer is 405
|
||||
// with an Allow header rather than a 404 that suggests the endpoint is
|
||||
// absent.
|
||||
mux.Handle("GET /mcp", server.Handler())
|
||||
|
||||
return 2
|
||||
}
|
||||
|
||||
// sessionResolver adapts the existing cookie session to oauth.SessionResolver.
|
||||
//
|
||||
// This is the ONLY place the OAuth package learns who is signed in, and it does
|
||||
// so through the existing session manager — the same lookup every other
|
||||
// authenticated route performs. No second password store, no second session
|
||||
// table, no second notion of identity.
|
||||
type sessionResolver struct{ s *Server }
|
||||
|
||||
// CurrentUser resolves the session cookie into an identity.
|
||||
//
|
||||
// Re-reads the user row rather than trusting the session's own copy, exactly as
|
||||
// authenticate() does, so a suspended account cannot approve an authorization
|
||||
// in the window before its session lapses.
|
||||
func (r sessionResolver) CurrentUser(req *http.Request) (authctx.Identity, bool) {
|
||||
token := sessionToken(req)
|
||||
if token == "" {
|
||||
return authctx.Identity{}, false
|
||||
}
|
||||
sess, err := r.s.sessions.Authenticate(req.Context(), token)
|
||||
if err != nil {
|
||||
return authctx.Identity{}, false
|
||||
}
|
||||
user, err := r.s.users.FindByID(req.Context(), sess.UserID)
|
||||
if err != nil || !user.IsActive() {
|
||||
return authctx.Identity{}, false
|
||||
}
|
||||
return authctx.Identity{
|
||||
UserID: user.ID, OrgID: user.OrgID, Email: user.Email,
|
||||
FullName: user.FullName, Role: user.Role, AccountType: user.AccountType,
|
||||
Status: user.Status, SessionID: sess.ID, ExpiresAt: sess.ExpiresAt,
|
||||
}, true
|
||||
}
|
||||
|
||||
// compile-time proof that the existing user store satisfies what OAuth needs.
|
||||
var _ oauth.UserLookup = (auth.UserStore)(nil)
|
||||
785
go-api/internal/httpserver/mcp_routes_test.go
Normal file
785
go-api/internal/httpserver/mcp_routes_test.go
Normal file
@@ -0,0 +1,785 @@
|
||||
package httpserver_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/config"
|
||||
"github.com/krow/krow-backend/go-api/internal/db"
|
||||
"github.com/krow/krow-backend/go-api/internal/httpserver"
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
// The configured deployment these tests run as. Fictional on purpose: every
|
||||
// URL in a discovery document must be traceable to THIS configuration, and a
|
||||
// realistic hostname would make a hardcoded one impossible to spot.
|
||||
const (
|
||||
testOAuthIssuer = "https://krow.example.test"
|
||||
testMCPResource = "https://krow.example.test/mcp"
|
||||
)
|
||||
|
||||
/* ── A fixture that can see headers and raw bodies ──────────────────────── */
|
||||
|
||||
// mcpResponse carries what the existing `response` deliberately does not: the
|
||||
// headers (WWW-Authenticate is the whole point of several tests) and the raw
|
||||
// body (the consent page is HTML, not JSON).
|
||||
//
|
||||
// A separate type rather than a change to `response`, so not one existing test
|
||||
// in this package is touched.
|
||||
type mcpResponse struct {
|
||||
code int
|
||||
body string
|
||||
header http.Header
|
||||
}
|
||||
|
||||
type mcpAPI struct {
|
||||
t *testing.T
|
||||
handler http.Handler
|
||||
srv *httpserver.Server
|
||||
h *testutil.Harness
|
||||
cookie *http.Cookie
|
||||
email string
|
||||
userID string
|
||||
}
|
||||
|
||||
// newOAuthAPI builds a server WITH OAuth configured, and signs in.
|
||||
//
|
||||
// The OAuth block is what makes routeOAuth and routeMCP register at all; the
|
||||
// standard newAPI fixture leaves it empty, which is what
|
||||
// TestMCPRoutesAreAbsentWhenUnconfigured relies on.
|
||||
func newOAuthAPI(t *testing.T) *mcpAPI {
|
||||
t.Helper()
|
||||
h := testutil.New(t)
|
||||
|
||||
cfg := &config.Config{
|
||||
AppEnv: "development",
|
||||
HTTP: config.HTTPConfig{
|
||||
Host: "127.0.0.1", Port: 0, ShutdownTimeout: time.Second,
|
||||
},
|
||||
DB: config.DBConfig{Schema: "public"},
|
||||
OAuth: config.OAuthConfig{
|
||||
Issuer: testOAuthIssuer,
|
||||
Resource: testMCPResource,
|
||||
LoginPath: "/login",
|
||||
},
|
||||
}
|
||||
log := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
srv, err := httpserver.New(cfg, &db.DB{Pool: h.Pool, Schema: "public"}, log)
|
||||
if err != nil {
|
||||
t.Fatalf("build the server: %v", err)
|
||||
}
|
||||
|
||||
a := &mcpAPI{t: t, handler: srv.Handler(), srv: srv, h: h}
|
||||
a.userID, a.email = seededUser(t, h.Pool)
|
||||
setPassword(t, h.Pool, a.userID)
|
||||
|
||||
result := signIn(t, a.handler, a.email, harnessPassword, false)
|
||||
if result.code != http.StatusOK || result.cookie == nil {
|
||||
t.Fatalf("the harness could not sign in: %d", result.code)
|
||||
}
|
||||
a.cookie = result.cookie
|
||||
return a
|
||||
}
|
||||
|
||||
func (a *mcpAPI) send(req *http.Request, withCookie bool) mcpResponse {
|
||||
a.t.Helper()
|
||||
if withCookie && a.cookie != nil {
|
||||
req.AddCookie(a.cookie)
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
a.handler.ServeHTTP(rec, req)
|
||||
return mcpResponse{code: rec.Code, body: rec.Body.String(), header: rec.Header()}
|
||||
}
|
||||
|
||||
func (a *mcpAPI) jsonReq(method, path string, payload any) *http.Request {
|
||||
a.t.Helper()
|
||||
var body io.Reader
|
||||
if payload != nil {
|
||||
raw, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
a.t.Fatalf("encode: %v", err)
|
||||
}
|
||||
body = bytes.NewReader(raw)
|
||||
}
|
||||
req := httptest.NewRequest(method, path, body)
|
||||
if payload != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
return req
|
||||
}
|
||||
|
||||
// do sends WITH the session cookie — a signed-in browser.
|
||||
func (a *mcpAPI) do(method, path string, payload any) mcpResponse {
|
||||
return a.send(a.jsonReq(method, path, payload), true)
|
||||
}
|
||||
|
||||
// doAnon sends WITHOUT any credential.
|
||||
func (a *mcpAPI) doAnon(method, path string, payload any) mcpResponse {
|
||||
return a.send(a.jsonReq(method, path, payload), false)
|
||||
}
|
||||
|
||||
// doAnonWithHeader sends one extra header and no cookie.
|
||||
func (a *mcpAPI) doAnonWithHeader(method, path string, payload any, key, value string) mcpResponse {
|
||||
req := a.jsonReq(method, path, payload)
|
||||
req.Header.Set(key, value)
|
||||
return a.send(req, false)
|
||||
}
|
||||
|
||||
func (a *mcpAPI) formReq(method, path string, form url.Values) *http.Request {
|
||||
req := httptest.NewRequest(method, path, strings.NewReader(form.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
return req
|
||||
}
|
||||
|
||||
// doForm posts a form WITH the cookie — the consent decision.
|
||||
func (a *mcpAPI) doForm(method, path string, form url.Values) mcpResponse {
|
||||
return a.send(a.formReq(method, path, form), true)
|
||||
}
|
||||
|
||||
// doAnonForm posts a form WITHOUT a cookie — the back-channel token call.
|
||||
func (a *mcpAPI) doAnonForm(method, path string, form url.Values) mcpResponse {
|
||||
return a.send(a.formReq(method, path, form), false)
|
||||
}
|
||||
|
||||
// oauthAccessToken runs the whole flow and returns a usable access token, for
|
||||
// tests that need a valid credential to prove it is being ignored.
|
||||
func (a *mcpAPI) oauthAccessToken(t *testing.T) string {
|
||||
t.Helper()
|
||||
reg := a.doAnon("POST", "/oauth/register", map[string]any{
|
||||
"client_name": "Token Helper", "redirect_uris": []string{"https://client.example.test/cb"},
|
||||
})
|
||||
var regDoc struct {
|
||||
ClientID string `json:"client_id"`
|
||||
}
|
||||
mustJSON(t, reg.body, ®Doc)
|
||||
|
||||
verifier := "helperVerifier0123456789abcdefghijklmnopqrst"
|
||||
q := url.Values{
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"response_type": {"code"}, "state": {"helper"},
|
||||
"code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"},
|
||||
"resource": {testMCPResource}, "scope": {"krow.read"},
|
||||
}
|
||||
consent := a.do("GET", "/oauth/authorize?"+q.Encode(), nil)
|
||||
csrf := between(consent.body, `name="csrf" value="`, `"`)
|
||||
|
||||
form := url.Values{}
|
||||
for k, v := range q {
|
||||
form[k] = v
|
||||
}
|
||||
form.Set("decision", "approve")
|
||||
form.Set("csrf", csrf)
|
||||
approved := a.doForm("POST", "/oauth/authorize", form)
|
||||
loc, _ := url.Parse(approved.header.Get("Location"))
|
||||
|
||||
tok := a.doAnonForm("POST", "/oauth/token", url.Values{
|
||||
"grant_type": {"authorization_code"}, "code": {loc.Query().Get("code")},
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"code_verifier": {verifier},
|
||||
})
|
||||
var tokens struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
}
|
||||
mustJSON(t, tok.body, &tokens)
|
||||
if tokens.AccessToken == "" {
|
||||
t.Fatalf("could not obtain a token: %s", tok.body)
|
||||
}
|
||||
return tokens.AccessToken
|
||||
}
|
||||
|
||||
// challengeFor derives an S256 challenge, so these tests do not depend on the
|
||||
// oauth package's unexported helpers.
|
||||
func challengeFor(verifier string) string {
|
||||
sum := sha256.Sum256([]byte(verifier))
|
||||
return base64.RawURLEncoding.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// The mounted surface, end to end.
|
||||
//
|
||||
// Everything below drives the REAL router — the same mux, the same
|
||||
// authenticate() middleware, the same publicPaths allowlist that serves
|
||||
// production. The point is not to re-test the OAuth package (internal/oauth
|
||||
// does that against its own handlers) but to prove the MOUNTING is right: that
|
||||
// discovery is reachable without a cookie, that /mcp is not, that a cookie
|
||||
// cannot substitute for a bearer token, and that the routes appear at all only
|
||||
// when the deployment is configured for them.
|
||||
|
||||
/* ── Route registration is conditional ──────────────────────────────────── */
|
||||
|
||||
// Without OAUTH_ISSUER and MCP_RESOURCE, none of this exists. An upgrade must
|
||||
// not quietly add an authorization server to a deployment that never asked.
|
||||
func TestMCPRoutesAreAbsentWhenUnconfigured(t *testing.T) {
|
||||
a := newAPI(t) // the standard fixture: no OAuth configuration
|
||||
|
||||
for _, path := range []string{
|
||||
"/mcp",
|
||||
"/oauth/register",
|
||||
"/oauth/authorize",
|
||||
"/oauth/token",
|
||||
"/.well-known/oauth-protected-resource",
|
||||
"/.well-known/oauth-authorization-server",
|
||||
} {
|
||||
r := a.doAnon("POST", path, nil)
|
||||
if r.code != http.StatusNotFound && r.code != http.StatusUnauthorized {
|
||||
t.Errorf("%s = %d on an unconfigured deployment; want 404 or 401, never a served response",
|
||||
path, r.code)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Discovery is public ────────────────────────────────────────────────── */
|
||||
|
||||
// A client with no token must be able to read both documents, or it can never
|
||||
// discover how to get one.
|
||||
func TestDiscoveryIsReachableWithoutASession(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
|
||||
t.Run("protected resource", func(t *testing.T) {
|
||||
r := a.doAnon("GET", "/.well-known/oauth-protected-resource", nil)
|
||||
if r.code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200 without a cookie: %s", r.code, r.body)
|
||||
}
|
||||
var doc struct {
|
||||
Resource string `json:"resource"`
|
||||
AuthorizationServers []string `json:"authorization_servers"`
|
||||
BearerMethods []string `json:"bearer_methods_supported"`
|
||||
}
|
||||
mustJSON(t, r.body, &doc)
|
||||
|
||||
if doc.Resource != testMCPResource {
|
||||
t.Errorf("resource = %q, want %q", doc.Resource, testMCPResource)
|
||||
}
|
||||
if len(doc.AuthorizationServers) != 1 || doc.AuthorizationServers[0] != testOAuthIssuer {
|
||||
t.Errorf("authorization_servers = %v, want [%q]", doc.AuthorizationServers, testOAuthIssuer)
|
||||
}
|
||||
// The MCP spec forbids a token in the query string.
|
||||
if strings.Join(doc.BearerMethods, ",") != "header" {
|
||||
t.Errorf("bearer_methods_supported = %v, want [header]", doc.BearerMethods)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("authorization server", func(t *testing.T) {
|
||||
r := a.doAnon("GET", "/.well-known/oauth-authorization-server", nil)
|
||||
if r.code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200 without a cookie: %s", r.code, r.body)
|
||||
}
|
||||
var doc struct {
|
||||
Issuer string `json:"issuer"`
|
||||
AuthorizationEndpoint string `json:"authorization_endpoint"`
|
||||
TokenEndpoint string `json:"token_endpoint"`
|
||||
RegistrationEndpoint string `json:"registration_endpoint"`
|
||||
Scopes []string `json:"scopes_supported"`
|
||||
ResponseTypes []string `json:"response_types_supported"`
|
||||
GrantTypes []string `json:"grant_types_supported"`
|
||||
PKCEMethods []string `json:"code_challenge_methods_supported"`
|
||||
ResourceIndicators bool `json:"resource_indicators_supported"`
|
||||
}
|
||||
mustJSON(t, r.body, &doc)
|
||||
|
||||
// EVERY url must come from configuration. A hardcoded hostname would
|
||||
// be one deployment's identity baked into every other one.
|
||||
if doc.Issuer != testOAuthIssuer {
|
||||
t.Errorf("issuer = %q, want %q", doc.Issuer, testOAuthIssuer)
|
||||
}
|
||||
for name, got := range map[string]string{
|
||||
"authorization_endpoint": doc.AuthorizationEndpoint,
|
||||
"token_endpoint": doc.TokenEndpoint,
|
||||
"registration_endpoint": doc.RegistrationEndpoint,
|
||||
} {
|
||||
if !strings.HasPrefix(got, testOAuthIssuer) {
|
||||
t.Errorf("%s = %q, want it under the configured issuer", name, got)
|
||||
}
|
||||
}
|
||||
if strings.Join(doc.ResponseTypes, ",") != "code" {
|
||||
t.Errorf("response_types_supported = %v; implicit must not be advertised", doc.ResponseTypes)
|
||||
}
|
||||
if strings.Join(doc.PKCEMethods, ",") != "S256" {
|
||||
t.Errorf("code_challenge_methods_supported = %v, want [S256]", doc.PKCEMethods)
|
||||
}
|
||||
for _, forbidden := range []string{"password", "client_credentials", "implicit"} {
|
||||
for _, advertised := range doc.GrantTypes {
|
||||
if advertised == forbidden {
|
||||
t.Errorf("grant_types_supported advertises %q", forbidden)
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, s := range doc.Scopes {
|
||||
if s == "krow.write" {
|
||||
t.Error("scopes_supported advertises krow.write")
|
||||
}
|
||||
}
|
||||
if !doc.ResourceIndicators {
|
||||
t.Error("resource_indicators_supported must be true")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/* ── /mcp authentication ────────────────────────────────────────────────── */
|
||||
|
||||
// No bearer → 401 with a challenge that tells the client where to go.
|
||||
func TestMCPWithoutBearerReturns401AndDiscoveryPointer(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
r := a.doAnon("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 1, "method": "tools/list",
|
||||
})
|
||||
|
||||
if r.code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401", r.code)
|
||||
}
|
||||
challenge := r.header.Get("WWW-Authenticate")
|
||||
if !strings.HasPrefix(challenge, "Bearer") {
|
||||
t.Fatalf("WWW-Authenticate = %q, want a Bearer challenge", challenge)
|
||||
}
|
||||
// RFC 9728: without resource_metadata the client has a 401 and nowhere to
|
||||
// look. This is the difference between "failed" and "here is how".
|
||||
if !strings.Contains(challenge, `resource_metadata="`+testOAuthIssuer) {
|
||||
t.Errorf("WWW-Authenticate = %q, want resource_metadata built from the configured issuer", challenge)
|
||||
}
|
||||
// And it must be built from config, not baked in.
|
||||
if strings.Contains(challenge, "krowforce.com") {
|
||||
t.Errorf("WWW-Authenticate contains a hardcoded production hostname: %q", challenge)
|
||||
}
|
||||
}
|
||||
|
||||
// THE test for this phase's riskiest decision: a perfectly valid KROW session
|
||||
// cookie must not open the MCP endpoint.
|
||||
func TestMCPRejectsACookieSession(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
|
||||
// `a.do` sends the authenticated session cookie the rest of the suite uses.
|
||||
r := a.do("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 1, "method": "tools/list",
|
||||
})
|
||||
|
||||
if r.code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401 — a browser cookie authenticated an MCP call", r.code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMCPRejectsAnInvalidBearer(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
for name, header := range map[string]string{
|
||||
"unknown token": "Bearer not-a-real-token",
|
||||
"empty": "Bearer ",
|
||||
"wrong scheme": "Basic dXNlcjpwYXNz",
|
||||
"no scheme": "abcdef",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
r := a.doAnonWithHeader("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 1, "method": "tools/list",
|
||||
}, "Authorization", header)
|
||||
if r.code != http.StatusUnauthorized {
|
||||
t.Errorf("status = %d, want 401", r.code)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// A token must never be accepted from the query string. The MCP spec forbids
|
||||
// it, and a URL is logged, cached and put in a Referer.
|
||||
func TestMCPIgnoresATokenInTheQueryString(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
token := a.oauthAccessToken(t)
|
||||
|
||||
r := a.doAnon("POST", "/mcp?access_token="+url.QueryEscape(token), map[string]any{
|
||||
"jsonrpc": "2.0", "id": 1, "method": "tools/list",
|
||||
})
|
||||
if r.code != http.StatusUnauthorized {
|
||||
t.Errorf("status = %d, want 401 — a query-string token was accepted", r.code)
|
||||
}
|
||||
}
|
||||
|
||||
// Custom identity headers must be ignored outright.
|
||||
func TestMCPIgnoresCustomIdentityHeaders(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
for _, header := range []string{"X-Access-Token", "X-Api-Key", "X-Org-Id", "X-User-Id", "X-Krow-Token"} {
|
||||
r := a.doAnonWithHeader("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 1, "method": "tools/list",
|
||||
}, header, a.oauthAccessToken(t))
|
||||
if r.code != http.StatusUnauthorized {
|
||||
t.Errorf("%s was accepted as a credential: %d", header, r.code)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── The full discovery → consent → token → MCP journey ─────────────────── */
|
||||
|
||||
// Every step a Claude client performs, over the real router, in order.
|
||||
func TestFullMCPConnectionJourney(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
|
||||
// 1–2. Call /mcp with no token; get 401 and a pointer.
|
||||
unauth := a.doAnon("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 1, "method": "initialize",
|
||||
})
|
||||
if unauth.code != http.StatusUnauthorized {
|
||||
t.Fatalf("step 1: status = %d, want 401", unauth.code)
|
||||
}
|
||||
challenge := unauth.header.Get("WWW-Authenticate")
|
||||
|
||||
// 3. Follow resource_metadata to the protected-resource document.
|
||||
metaURL := between(challenge, `resource_metadata="`, `"`)
|
||||
if metaURL == "" {
|
||||
t.Fatal("step 3: the challenge carries no resource_metadata")
|
||||
}
|
||||
prPath := strings.TrimPrefix(metaURL, testOAuthIssuer)
|
||||
pr := a.doAnon("GET", prPath, nil)
|
||||
if pr.code != http.StatusOK {
|
||||
t.Fatalf("step 3: %s = %d", prPath, pr.code)
|
||||
}
|
||||
var prDoc struct {
|
||||
AuthorizationServers []string `json:"authorization_servers"`
|
||||
}
|
||||
mustJSON(t, pr.body, &prDoc)
|
||||
|
||||
// 4. Authorization-server metadata.
|
||||
as := a.doAnon("GET", "/.well-known/oauth-authorization-server", nil)
|
||||
if as.code != http.StatusOK {
|
||||
t.Fatalf("step 4: status = %d", as.code)
|
||||
}
|
||||
var asDoc struct {
|
||||
AuthorizationEndpoint string `json:"authorization_endpoint"`
|
||||
TokenEndpoint string `json:"token_endpoint"`
|
||||
RegistrationEndpoint string `json:"registration_endpoint"`
|
||||
}
|
||||
mustJSON(t, as.body, &asDoc)
|
||||
|
||||
// 5. Register, at the advertised endpoint.
|
||||
reg := a.doAnon("POST", strings.TrimPrefix(asDoc.RegistrationEndpoint, testOAuthIssuer), map[string]any{
|
||||
"client_name": "Journey Client",
|
||||
"redirect_uris": []string{"https://client.example.test/cb"},
|
||||
})
|
||||
if reg.code != http.StatusCreated {
|
||||
t.Fatalf("step 5: registration = %d %s", reg.code, reg.body)
|
||||
}
|
||||
var regDoc struct {
|
||||
ClientID string `json:"client_id"`
|
||||
}
|
||||
mustJSON(t, reg.body, ®Doc)
|
||||
|
||||
// 6–7. Authorize, SIGNED IN. A cookie is exactly right here: this step is
|
||||
// a person in a browser.
|
||||
verifier := "journeyVerifier0123456789abcdefghijklmnopqrs"
|
||||
q := url.Values{
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"response_type": {"code"}, "state": {"journey-state"},
|
||||
"code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"},
|
||||
"resource": {testMCPResource}, "scope": {"krow.read"},
|
||||
}
|
||||
consent := a.do("GET", "/oauth/authorize?"+q.Encode(), nil)
|
||||
|
||||
// 8. A consent page, not a code.
|
||||
if consent.code != http.StatusOK {
|
||||
t.Fatalf("step 8: expected a consent page, got %d %s", consent.code, consent.body)
|
||||
}
|
||||
if !strings.Contains(consent.body, "Journey Client") {
|
||||
t.Error("step 8: the consent page does not name the requesting client")
|
||||
}
|
||||
csrf := between(consent.body, `name="csrf" value="`, `"`)
|
||||
if csrf == "" {
|
||||
t.Fatal("step 8: no csrf token in the consent form")
|
||||
}
|
||||
|
||||
// 9–10. Approve; receive a code.
|
||||
form := url.Values{}
|
||||
for k, v := range q {
|
||||
form[k] = v
|
||||
}
|
||||
form.Set("decision", "approve")
|
||||
form.Set("csrf", csrf)
|
||||
approved := a.doForm("POST", "/oauth/authorize", form)
|
||||
if approved.code != http.StatusFound {
|
||||
t.Fatalf("step 10: approve = %d %s", approved.code, approved.body)
|
||||
}
|
||||
loc, _ := url.Parse(approved.header.Get("Location"))
|
||||
code := loc.Query().Get("code")
|
||||
if code == "" {
|
||||
t.Fatalf("step 10: no code in %s", loc)
|
||||
}
|
||||
if loc.Query().Get("state") != "journey-state" {
|
||||
t.Errorf("step 10: state = %q", loc.Query().Get("state"))
|
||||
}
|
||||
|
||||
// 11. Exchange — with NO cookie, as a back-channel call.
|
||||
tok := a.doAnonForm("POST", strings.TrimPrefix(asDoc.TokenEndpoint, testOAuthIssuer), url.Values{
|
||||
"grant_type": {"authorization_code"}, "code": {code},
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"code_verifier": {verifier},
|
||||
})
|
||||
if tok.code != http.StatusOK {
|
||||
t.Fatalf("step 11: token = %d %s", tok.code, tok.body)
|
||||
}
|
||||
var tokens struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
TokenType string `json:"token_type"`
|
||||
}
|
||||
mustJSON(t, tok.body, &tokens)
|
||||
if tokens.AccessToken == "" || tokens.TokenType != "Bearer" {
|
||||
t.Fatalf("step 11: unusable token response: %s", tok.body)
|
||||
}
|
||||
|
||||
// 12–13. tools/list with the bearer token.
|
||||
list := a.doAnonWithHeader("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 2, "method": "tools/list",
|
||||
}, "Authorization", "Bearer "+tokens.AccessToken)
|
||||
if list.code != http.StatusOK {
|
||||
t.Fatalf("step 13: tools/list = %d %s", list.code, list.body)
|
||||
}
|
||||
var listDoc struct {
|
||||
Result struct {
|
||||
Tools []struct {
|
||||
Name string `json:"name"`
|
||||
} `json:"tools"`
|
||||
} `json:"result"`
|
||||
}
|
||||
mustJSON(t, list.body, &listDoc)
|
||||
if len(listDoc.Result.Tools) != 16 {
|
||||
t.Errorf("step 13: %d tools, want 16", len(listDoc.Result.Tools))
|
||||
}
|
||||
for _, tool := range listDoc.Result.Tools {
|
||||
switch tool.Name {
|
||||
case "assign_worker", "move_application", "knowledge_search":
|
||||
t.Errorf("step 13: %q is exposed over the mounted route", tool.Name)
|
||||
}
|
||||
}
|
||||
|
||||
// 14. tools/call reaches the existing authorization and real data.
|
||||
call := a.doAnonWithHeader("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 3, "method": "tools/call",
|
||||
"params": map[string]any{"name": "workspace_summary", "arguments": map[string]any{}},
|
||||
}, "Authorization", "Bearer "+tokens.AccessToken)
|
||||
if call.code != http.StatusOK {
|
||||
t.Fatalf("step 14: tools/call = %d %s", call.code, call.body)
|
||||
}
|
||||
var callDoc struct {
|
||||
Result struct {
|
||||
IsError bool `json:"isError"`
|
||||
Content []struct {
|
||||
Text string `json:"text"`
|
||||
} `json:"content"`
|
||||
} `json:"result"`
|
||||
}
|
||||
mustJSON(t, call.body, &callDoc)
|
||||
if callDoc.Result.IsError {
|
||||
t.Fatalf("step 14: the tool refused: %s", callDoc.Result.Content[0].Text)
|
||||
}
|
||||
}
|
||||
|
||||
// Denial must reach the client correctly and issue nothing.
|
||||
func TestConsentDenialOverTheMountedRoute(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
|
||||
reg := a.doAnon("POST", "/oauth/register", map[string]any{
|
||||
"client_name": "Deny Client", "redirect_uris": []string{"https://client.example.test/cb"},
|
||||
})
|
||||
var regDoc struct {
|
||||
ClientID string `json:"client_id"`
|
||||
}
|
||||
mustJSON(t, reg.body, ®Doc)
|
||||
|
||||
verifier := "denyVerifier0123456789abcdefghijklmnopqrstuv"
|
||||
q := url.Values{
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"response_type": {"code"}, "state": {"deny-state"},
|
||||
"code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"},
|
||||
"resource": {testMCPResource}, "scope": {"krow.read"},
|
||||
}
|
||||
consent := a.do("GET", "/oauth/authorize?"+q.Encode(), nil)
|
||||
csrf := between(consent.body, `name="csrf" value="`, `"`)
|
||||
|
||||
form := url.Values{}
|
||||
for k, v := range q {
|
||||
form[k] = v
|
||||
}
|
||||
form.Set("decision", "deny")
|
||||
form.Set("csrf", csrf)
|
||||
|
||||
denied := a.doForm("POST", "/oauth/authorize", form)
|
||||
if denied.code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want 302", denied.code)
|
||||
}
|
||||
loc, _ := url.Parse(denied.header.Get("Location"))
|
||||
if got := loc.Query().Get("error"); got != "access_denied" {
|
||||
t.Errorf("error = %q, want access_denied", got)
|
||||
}
|
||||
if got := loc.Query().Get("state"); got != "deny-state" {
|
||||
t.Errorf("state = %q, want deny-state", got)
|
||||
}
|
||||
if loc.Query().Get("code") != "" {
|
||||
t.Error("a denial issued a code")
|
||||
}
|
||||
}
|
||||
|
||||
// /oauth/authorize is NOT public: an anonymous visitor must be sent to login.
|
||||
func TestAuthorizeRequiresASession(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
r := a.doAnon("GET", "/oauth/authorize?client_id=x", nil)
|
||||
|
||||
// Either the middleware refuses it (401) or the handler redirects to
|
||||
// login. Both are correct; serving a consent page is not.
|
||||
if r.code == http.StatusOK && strings.Contains(r.body, "Approve") {
|
||||
t.Fatal("a consent page was served to an anonymous visitor")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Helpers ────────────────────────────────────────────────────────────── */
|
||||
|
||||
func mustJSON(t *testing.T, body string, dst any) {
|
||||
t.Helper()
|
||||
if err := json.Unmarshal([]byte(body), dst); err != nil {
|
||||
t.Fatalf("response was not JSON: %v\nbody: %s", err, body)
|
||||
}
|
||||
}
|
||||
|
||||
// between returns the text between two markers, or "".
|
||||
func between(s, start, end string) string {
|
||||
i := strings.Index(s, start)
|
||||
if i < 0 {
|
||||
return ""
|
||||
}
|
||||
rest := s[i+len(start):]
|
||||
j := strings.Index(rest, end)
|
||||
if j < 0 {
|
||||
return ""
|
||||
}
|
||||
return rest[:j]
|
||||
}
|
||||
|
||||
/* ── Anonymous /oauth/authorize must reach the handler ──────────────────── */
|
||||
|
||||
// The regression test for the defect a live Claude Web connection exposed.
|
||||
//
|
||||
// /oauth/authorize was withheld from publicPaths, so the cookie middleware
|
||||
// answered a signed-out visitor with its JSON 401 and the handler never ran —
|
||||
// which meant the handler's redirect-to-login could never execute. A first-time
|
||||
// connector user is signed out by definition, so OAuth's browser leg was
|
||||
// unreachable for precisely the people who needed it.
|
||||
//
|
||||
// WHY THE EXISTING TESTS MISSED IT, and why this one is shaped differently:
|
||||
//
|
||||
// - oauth.TestAuthorizeRedirectsAnonymousToLogin drives AuthorizeHandler
|
||||
// DIRECTLY, so the middleware is not in the path at all. It passed against
|
||||
// broken behaviour because it never exercised the thing that was broken.
|
||||
// - TestAuthorizeRequiresASession (below) asserts only that a consent page is
|
||||
// not served anonymously — which a 401 satisfies perfectly well.
|
||||
//
|
||||
// So this one drives the MOUNTED router and asserts the POSITIVE behaviour: a
|
||||
// redirect to the login, carrying the original authorization request.
|
||||
func TestAnonymousAuthorizeReachesTheHandlerAndRedirectsToLogin(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
|
||||
// A client to name, so the request is well-formed enough to get past the
|
||||
// handler's own client/redirect validation and reach the session check.
|
||||
reg := a.doAnon("POST", "/oauth/register", map[string]any{
|
||||
"client_name": "Anonymous Flow", "redirect_uris": []string{"https://client.example.test/cb"},
|
||||
})
|
||||
var regDoc struct {
|
||||
ClientID string `json:"client_id"`
|
||||
}
|
||||
mustJSON(t, reg.body, ®Doc)
|
||||
|
||||
verifier := "anonVerifier0123456789abcdefghijklmnopqrstu"
|
||||
q := url.Values{
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"response_type": {"code"}, "state": {"anon-state"},
|
||||
"code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"},
|
||||
"resource": {testMCPResource}, "scope": {"krow.read"},
|
||||
}
|
||||
|
||||
// doAnon sends NO session cookie — a first-time connector user.
|
||||
r := a.doAnon("GET", "/oauth/authorize?"+q.Encode(), nil)
|
||||
|
||||
// The defect: the middleware's JSON 401 instead of the handler's redirect.
|
||||
if r.code == http.StatusUnauthorized {
|
||||
t.Fatalf("the middleware refused before the handler ran: %d %s\n"+
|
||||
"a signed-out visitor must be sent to sign in, not told 'no'", r.code, r.body)
|
||||
}
|
||||
if strings.Contains(r.body, `"code": "unauthorized"`) ||
|
||||
strings.Contains(r.body, `"code":"unauthorized"`) {
|
||||
t.Fatalf("the response is the middleware's JSON 401, not the handler's: %s", r.body)
|
||||
}
|
||||
|
||||
if r.code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want 302 to the login", r.code)
|
||||
}
|
||||
location := r.header.Get("Location")
|
||||
if !strings.HasPrefix(location, "/login?returnTo=") {
|
||||
t.Fatalf("Location = %q, want a redirect to the configured login path", location)
|
||||
}
|
||||
|
||||
// The whole authorization request must survive the round trip, or the
|
||||
// person signs in and lands nowhere.
|
||||
returnTo, err := url.QueryUnescape(strings.TrimPrefix(location, "/login?returnTo="))
|
||||
if err != nil {
|
||||
t.Fatalf("returnTo is not decodable: %v", err)
|
||||
}
|
||||
for name, want := range map[string]string{
|
||||
"path": "/oauth/authorize",
|
||||
"client_id": "client_id=" + regDoc.ClientID,
|
||||
"state": "state=anon-state",
|
||||
"code_challenge": "code_challenge=" + challengeFor(verifier),
|
||||
"code_challenge_method": "code_challenge_method=S256",
|
||||
"resource": "resource=",
|
||||
"redirect_uri": "redirect_uri=",
|
||||
} {
|
||||
if !strings.Contains(returnTo, want) {
|
||||
t.Errorf("returnTo has lost the %s: %q", name, returnTo)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Listing the path must NOT hand out consent, or a code, to somebody signed
|
||||
// out. "Public" here means the handler decides — not that the route is open.
|
||||
func TestAnonymousAuthorizeStillGrantsNothing(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
|
||||
reg := a.doAnon("POST", "/oauth/register", map[string]any{
|
||||
"client_name": "Nothing Granted", "redirect_uris": []string{"https://client.example.test/cb"},
|
||||
})
|
||||
var regDoc struct {
|
||||
ClientID string `json:"client_id"`
|
||||
}
|
||||
mustJSON(t, reg.body, ®Doc)
|
||||
|
||||
verifier := "nothingVerifier0123456789abcdefghijklmnopq"
|
||||
q := url.Values{
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"response_type": {"code"}, "state": {"nothing"},
|
||||
"code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"},
|
||||
"resource": {testMCPResource}, "scope": {"krow.read"},
|
||||
}
|
||||
|
||||
// A GET must not render consent.
|
||||
get := a.doAnon("GET", "/oauth/authorize?"+q.Encode(), nil)
|
||||
if strings.Contains(get.body, "Approve") || strings.Contains(get.body, "Authorize access to Krow") {
|
||||
t.Error("a consent page was served to a signed-out visitor")
|
||||
}
|
||||
|
||||
// And a POST — skipping the page entirely, as an attacker would — must not
|
||||
// issue a code. The handler's session check refuses before the CSRF check
|
||||
// is even relevant.
|
||||
form := url.Values{}
|
||||
for k, v := range q {
|
||||
form[k] = v
|
||||
}
|
||||
form.Set("decision", "approve")
|
||||
form.Set("csrf", "forged")
|
||||
|
||||
post := a.doAnonForm("POST", "/oauth/authorize", form)
|
||||
if loc := post.header.Get("Location"); strings.Contains(loc, "code=") {
|
||||
t.Fatalf("an anonymous POST obtained an authorization code: %s", loc)
|
||||
}
|
||||
if post.code == http.StatusFound && strings.HasPrefix(post.header.Get("Location"), "https://client.example.test") {
|
||||
t.Fatalf("an anonymous POST reached the client callback: %s", post.header.Get("Location"))
|
||||
}
|
||||
}
|
||||
230
go-api/internal/httpserver/mcplimit.go
Normal file
230
go-api/internal/httpserver/mcplimit.go
Normal file
@@ -0,0 +1,230 @@
|
||||
package httpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/ratelimit"
|
||||
)
|
||||
|
||||
// Rate limiting for the mounted OAuth and MCP routes.
|
||||
//
|
||||
// A middleware rather than a change inside either package, for one reason: the
|
||||
// SUBJECT of a limit is an HTTP concept. Which IP, which bearer token, which
|
||||
// form field names the client — none of that is knowable from inside
|
||||
// internal/oauth, and handing those packages a request so they could work it
|
||||
// out would put transport details in a layer that has none.
|
||||
//
|
||||
// NOTHING SENSITIVE BECOMES A BUCKET KEY. Every subject goes through
|
||||
// ratelimit.Subject, which hashes it. A bucket naming a bearer token would
|
||||
// write that token into a table and into any log line mentioning the bucket.
|
||||
|
||||
// limited wraps a handler with one rule, keyed by a subject derived per request.
|
||||
//
|
||||
// The subject function returns "" to mean "not limitable" — no token on the
|
||||
// request, say — and the request passes through. That is correct rather than
|
||||
// lax: a request with no identifiable subject is refused by the handler itself
|
||||
// a moment later, and inventing a shared bucket for all of them would let one
|
||||
// caller exhaust a budget that everybody else then queues behind.
|
||||
func (s *Server) limited(rule ratelimit.Rule, subject func(*http.Request) string, next http.Handler) http.Handler {
|
||||
if s.limiter == nil {
|
||||
return next
|
||||
}
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
raw := subject(r)
|
||||
if raw == "" {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
decision, err := s.limiter.Allow(r.Context(), rule, ratelimit.Subject(raw))
|
||||
if err != nil {
|
||||
// The limiter could not count. It has already decided whether that
|
||||
// permits the request — fail closed by default — and this logs the
|
||||
// fault without naming the subject, which is a hash of a
|
||||
// credential.
|
||||
s.log.Error("rate limiter unavailable",
|
||||
"rule", rule.Name, "allowed", decision.Allowed, "error", err)
|
||||
}
|
||||
|
||||
// Headers on every response, not only refusals, so a well-behaved
|
||||
// client can slow down before it is refused rather than after.
|
||||
w.Header().Set("RateLimit-Limit", strconv.Itoa(decision.Limit))
|
||||
w.Header().Set("RateLimit-Remaining", strconv.Itoa(decision.Remaining))
|
||||
|
||||
if !decision.Allowed {
|
||||
// Retry-After in seconds, rounded up and never zero — "Retry-After:
|
||||
// 0" invites an immediate retry, which is the one thing a limited
|
||||
// client must not do. The same reasoning as retryAfterSeconds in
|
||||
// ratelimit.go, and the same rounding.
|
||||
w.Header().Set("Retry-After", retryAfterSeconds(decision.RetryAfter))
|
||||
s.log.Warn("rate limit exceeded", "rule", rule.Name, "path", r.URL.Path)
|
||||
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
w.WriteHeader(http.StatusTooManyRequests)
|
||||
_, _ = w.Write([]byte(`{"error":"rate_limited",` +
|
||||
`"error_description":"too many requests; retry after the interval in the Retry-After header"}`))
|
||||
return
|
||||
}
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
/* ── Subjects ───────────────────────────────────────────────────────────── */
|
||||
|
||||
// byClientAddr keys by the caller's address, for endpoints with no credential.
|
||||
//
|
||||
// A method, not a free function, because the address is no longer a property of
|
||||
// the request alone: resolving it needs the trusted-proxy set, which is wired
|
||||
// onto the server. See clientip.go.
|
||||
func (s *Server) byClientAddr(r *http.Request) string { return s.trust.clientAddr(r) }
|
||||
|
||||
// byAddrAndUser keys the authorization endpoint by address AND signed-in user.
|
||||
//
|
||||
// Both, because either alone is wrong: by user only, an attacker could exhaust
|
||||
// somebody else's budget by naming them; by address only, an office behind one
|
||||
// NAT shares one person's allowance.
|
||||
func (s *Server) byAddrAndUser(r *http.Request) string {
|
||||
subject := s.trust.clientAddr(r)
|
||||
if identity, ok := (sessionResolver{s}).CurrentUser(r); ok {
|
||||
subject += "|" + identity.UserID
|
||||
}
|
||||
return subject
|
||||
}
|
||||
|
||||
// byFormClientID keys the token endpoint by the client_id it names.
|
||||
//
|
||||
// Reading a form value means parsing the body, which the handler then parses
|
||||
// again — ParseForm caches on the request, so the second call is free.
|
||||
func (s *Server) byFormClientID(r *http.Request) string {
|
||||
if err := r.ParseForm(); err != nil {
|
||||
return ""
|
||||
}
|
||||
if id := r.PostFormValue("client_id"); id != "" {
|
||||
return id
|
||||
}
|
||||
// No client_id: the handler will refuse it. Fall back to the address so a
|
||||
// caller cannot dodge the limit by omitting the field.
|
||||
return s.trust.clientAddr(r)
|
||||
}
|
||||
|
||||
// byRefreshFamily keys refresh by the token being presented.
|
||||
//
|
||||
// Keyed by the TOKEN's hash, not the family id, because the family is not
|
||||
// knowable without a database read this middleware has no business doing. The
|
||||
// effect is very nearly the same: a rotation produces a new token and therefore
|
||||
// a new bucket, so the practical limit is per-token-per-window rather than
|
||||
// per-family — which bounds a loop just as well, since a loop presenting the
|
||||
// SAME token is exactly what the limit is for.
|
||||
func byRefreshToken(r *http.Request) string {
|
||||
if err := r.ParseForm(); err != nil {
|
||||
return ""
|
||||
}
|
||||
return r.PostFormValue("refresh_token")
|
||||
}
|
||||
|
||||
// byBearerToken keys MCP by the presented access token.
|
||||
//
|
||||
// The narrowest identity available on an MCP request, and the right one: it is
|
||||
// one connection from one client for one user. Keying by user would let a
|
||||
// person's second client eat the first's budget; keying by IP would make
|
||||
// Claude's shared egress one bucket for every customer.
|
||||
func byBearerToken(r *http.Request) string {
|
||||
const prefix = "Bearer "
|
||||
header := r.Header.Get("Authorization")
|
||||
if len(header) <= len(prefix) {
|
||||
return ""
|
||||
}
|
||||
// Case-insensitive prefix, matching mcpserver's own parsing.
|
||||
if !equalFoldASCII(header[:len(prefix)], prefix) {
|
||||
return ""
|
||||
}
|
||||
return header[len(prefix):]
|
||||
}
|
||||
|
||||
func equalFoldASCII(a, b string) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
for i := 0; i < len(a); i++ {
|
||||
ca, cb := a[i], b[i]
|
||||
if 'A' <= ca && ca <= 'Z' {
|
||||
ca += 'a' - 'A'
|
||||
}
|
||||
if 'A' <= cb && cb <= 'Z' {
|
||||
cb += 'a' - 'A'
|
||||
}
|
||||
if ca != cb {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// mcpLimited applies BOTH tool-call limits to the MCP endpoint.
|
||||
//
|
||||
// Two rules stacked rather than one, because they stop different things: the
|
||||
// per-minute rule bounds a spike, and the per-hour rule bounds a slow drain
|
||||
// that would sit under the per-minute rule indefinitely. Checked minute-first
|
||||
// so the cheaper refusal happens earlier.
|
||||
func (s *Server) mcpLimited(next http.Handler) http.Handler {
|
||||
return s.limited(ratelimit.MCPToolCallPerMinute, byBearerToken,
|
||||
s.limited(ratelimit.MCPToolCallPerHour, byBearerToken, next))
|
||||
}
|
||||
|
||||
// tokenLimited applies the right rule for the grant type being requested.
|
||||
//
|
||||
// One endpoint, two grant types, two different abuse shapes — so one limit
|
||||
// keyed one way would be wrong for the other. A code exchange is bounded per
|
||||
// client; a refresh is bounded per presented token, so one connection looping
|
||||
// cannot spend a second connection's budget.
|
||||
func (s *Server) tokenLimited(next http.Handler) http.Handler {
|
||||
exchange := s.limited(ratelimit.OAuthToken, s.byFormClientID, next)
|
||||
refresh := s.limited(ratelimit.OAuthRefresh, byRefreshToken, next)
|
||||
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
// ParseForm caches on the request, so the handler's own call is free.
|
||||
if err := r.ParseForm(); err != nil {
|
||||
next.ServeHTTP(w, r) // let the handler produce the proper error
|
||||
return
|
||||
}
|
||||
if r.PostFormValue("grant_type") == "refresh_token" {
|
||||
refresh.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
exchange.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
/* ── The per-organisation ceiling ───────────────────────────────────────── */
|
||||
|
||||
// orgLimiter adapts the shared limiter to mcpserver.OrgLimiter.
|
||||
//
|
||||
// It is the ONE limit that cannot live in the middleware above, because an
|
||||
// organisation is not knowable until the bearer token has been resolved to a
|
||||
// user and that user's row read. A middleware running before authentication
|
||||
// could only key by something the client supplied — which is exactly the
|
||||
// identity this surface refuses to trust.
|
||||
//
|
||||
// So mcpserver calls this from inside dispatch, after it has an
|
||||
// authctx.Identity, and passes the org id from that identity. This type has no
|
||||
// access to the request and therefore no way to be handed a different one.
|
||||
type orgLimiter struct{ s *Server }
|
||||
|
||||
// AllowOrg counts one call against the organisation's hourly ceiling.
|
||||
//
|
||||
// The org id is hashed like every other subject. It is not a secret, but the
|
||||
// bucket format is uniform and a uuid in a table of counters is one more place
|
||||
// a tenant identifier exists for no reason.
|
||||
func (o orgLimiter) AllowOrg(ctx context.Context, orgID string) (bool, time.Duration, error) {
|
||||
if o.s.limiter == nil || orgID == "" {
|
||||
// No limiter, or no organisation — the latter cannot happen, because
|
||||
// mcpserver refuses an identity without one before it reaches here.
|
||||
return true, 0, nil
|
||||
}
|
||||
d, err := o.s.limiter.Allow(ctx, ratelimit.MCPPerOrgPerHour, ratelimit.Subject(orgID))
|
||||
return d.Allowed, d.RetryAfter, err
|
||||
}
|
||||
294
go-api/internal/httpserver/proxybuckets_test.go
Normal file
294
go-api/internal/httpserver/proxybuckets_test.go
Normal file
@@ -0,0 +1,294 @@
|
||||
package httpserver_test
|
||||
|
||||
// Multi-client rate-limit isolation, end to end.
|
||||
//
|
||||
// WHAT THIS IS FOR
|
||||
//
|
||||
// clientip_test.go proves the address RESOLVER picks the right string. It does
|
||||
// not prove the string reaches Postgres as a distinct bucket, that the mounted
|
||||
// route uses the resolver at all, or that the OAuth registration limit is
|
||||
// actually per-client once it does. Those are different failures — a correct
|
||||
// resolver wired to nothing looks identical from a unit test — and they are
|
||||
// what broke in production, so they are tested here against the real handler
|
||||
// and the real limiter.
|
||||
//
|
||||
// Every test drives httpserver.Handler() through the full middleware stack with
|
||||
// a real database behind it. Nothing is stubbed.
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/config"
|
||||
"github.com/krow/krow-backend/go-api/internal/db"
|
||||
"github.com/krow/krow-backend/go-api/internal/httpserver"
|
||||
"github.com/krow/krow-backend/go-api/internal/ratelimit"
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
// The proxy this fixture's deployment sits behind, and an address inside it.
|
||||
const (
|
||||
proxyNetwork = "10.42.0.0/16"
|
||||
proxyAddr = "10.42.0.1:9999"
|
||||
)
|
||||
|
||||
// proxiedAPI is newOAuthAPI with a trusted proxy configured.
|
||||
//
|
||||
// Deliberately not a flag on newOAuthAPI: every existing test in this package
|
||||
// must keep running with an EMPTY trusted set, because that is the default
|
||||
// posture and a regression in it is the thing worth catching.
|
||||
type proxiedAPI struct {
|
||||
handler http.Handler
|
||||
h *testutil.Harness
|
||||
}
|
||||
|
||||
func newProxiedAPI(t *testing.T) *proxiedAPI {
|
||||
t.Helper()
|
||||
h := testutil.New(t)
|
||||
|
||||
network, err := netip.ParsePrefix(proxyNetwork)
|
||||
if err != nil {
|
||||
t.Fatalf("bad test network: %v", err)
|
||||
}
|
||||
|
||||
cfg := &config.Config{
|
||||
AppEnv: "development",
|
||||
HTTP: config.HTTPConfig{
|
||||
Host: "127.0.0.1", Port: 0, ShutdownTimeout: time.Second,
|
||||
TrustedProxies: []netip.Prefix{network.Masked()},
|
||||
},
|
||||
DB: config.DBConfig{Schema: "public"},
|
||||
OAuth: config.OAuthConfig{
|
||||
Issuer: testOAuthIssuer,
|
||||
Resource: testMCPResource,
|
||||
LoginPath: "/login",
|
||||
},
|
||||
}
|
||||
srv, err := httpserver.New(cfg, &db.DB{Pool: h.Pool, Schema: "public"},
|
||||
slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
if err != nil {
|
||||
t.Fatalf("build the server: %v", err)
|
||||
}
|
||||
return &proxiedAPI{handler: srv.Handler(), h: h}
|
||||
}
|
||||
|
||||
// register performs one DCR as a caller arriving via the proxy.
|
||||
//
|
||||
// peer is what net/http would report as RemoteAddr; forwarded is the
|
||||
// X-Forwarded-For the proxy appended. An empty forwarded value sends no header.
|
||||
func (a *proxiedAPI) register(t *testing.T, peer, forwarded, name string) int {
|
||||
t.Helper()
|
||||
body := fmt.Sprintf(
|
||||
`{"client_name":%q,"redirect_uris":["https://client.example.test/cb"]}`, name)
|
||||
req, err := http.NewRequest("POST", "/oauth/register", strings.NewReader(body))
|
||||
if err != nil {
|
||||
t.Fatalf("build request: %v", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.RemoteAddr = peer
|
||||
if forwarded != "" {
|
||||
req.Header.Set("X-Forwarded-For", forwarded)
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
a.handler.ServeHTTP(rec, req)
|
||||
return rec.Code
|
||||
}
|
||||
|
||||
// exhaust registers until the limit refuses, and returns how many succeeded.
|
||||
// It stops well past the limit so a failure reports a number rather than hanging.
|
||||
func (a *proxiedAPI) exhaust(t *testing.T, peer, forwarded, name string) int {
|
||||
t.Helper()
|
||||
allowed := 0
|
||||
for i := 0; i < ratelimit.OAuthRegister.Limit*3; i++ {
|
||||
code := a.register(t, peer, forwarded, fmt.Sprintf("%s-%d", name, i))
|
||||
if code == http.StatusTooManyRequests {
|
||||
return allowed
|
||||
}
|
||||
if code != http.StatusCreated {
|
||||
t.Fatalf("%s attempt %d: unexpected status %d", name, i, code)
|
||||
}
|
||||
allowed++
|
||||
}
|
||||
t.Fatalf("%s was never refused after %d registrations; the limit is not applied",
|
||||
name, ratelimit.OAuthRegister.Limit*3)
|
||||
return allowed
|
||||
}
|
||||
|
||||
/* ── Ten independent clients ────────────────────────────────────────────── */
|
||||
|
||||
// The headline requirement: many users behind one proxy each get their own
|
||||
// budget. Before this change every one of these shared a bucket and the
|
||||
// eleventh registration on the list would have been refused.
|
||||
func TestTenClientsBehindOneProxyDoNotShareABucket(t *testing.T) {
|
||||
a := newProxiedAPI(t)
|
||||
|
||||
const clients = 10
|
||||
for i := 0; i < clients; i++ {
|
||||
client := fmt.Sprintf("203.0.113.%d", i+1)
|
||||
// Each client registers TWICE — Claude mints a new client per connect,
|
||||
// so a reconnect must not count against anybody else.
|
||||
for attempt := 0; attempt < 2; attempt++ {
|
||||
code := a.register(t, proxyAddr, client+", "+"10.42.0.1",
|
||||
fmt.Sprintf("client-%d-%d", i, attempt))
|
||||
if code != http.StatusCreated {
|
||||
t.Fatalf("client %s attempt %d: status %d, want 201 — clients are sharing a bucket",
|
||||
client, attempt, code)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Twenty registrations went through on a limit of ten per subject, which is
|
||||
// only possible if the subject really is the forwarded client.
|
||||
var buckets int
|
||||
if err := a.h.Pool.QueryRow(t.Context(),
|
||||
`SELECT count(*) FROM rate_limits WHERE bucket LIKE 'oauth.register:%'`).Scan(&buckets); err != nil {
|
||||
t.Fatalf("count buckets: %v", err)
|
||||
}
|
||||
if buckets != clients {
|
||||
t.Errorf("%d distinct oauth.register buckets, want %d — one per client", buckets, clients)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── One client's exhaustion is its own ─────────────────────────────────── */
|
||||
|
||||
// Client A burns its whole budget; client B is unaffected. This is the property
|
||||
// that failed in production, where A's retries refused B outright.
|
||||
func TestOneClientExhaustingDoesNotBlockAnother(t *testing.T) {
|
||||
a := newProxiedAPI(t)
|
||||
|
||||
const clientA, clientB = "203.0.113.50", "203.0.113.51"
|
||||
|
||||
allowed := a.exhaust(t, proxyAddr, clientA, "A")
|
||||
if allowed != ratelimit.OAuthRegister.Limit {
|
||||
t.Errorf("client A got %d registrations, want %d", allowed, ratelimit.OAuthRegister.Limit)
|
||||
}
|
||||
|
||||
// A is now refused.
|
||||
if code := a.register(t, proxyAddr, clientA, "A-again"); code != http.StatusTooManyRequests {
|
||||
t.Errorf("client A after exhausting: status %d, want 429", code)
|
||||
}
|
||||
// B is not.
|
||||
if code := a.register(t, proxyAddr, clientB, "B"); code != http.StatusCreated {
|
||||
t.Errorf("client B: status %d, want 201 — A's exhaustion blocked B", code)
|
||||
}
|
||||
}
|
||||
|
||||
// A single client reconnecting repeatedly spends only its own budget, which is
|
||||
// what Claude actually does: a new DCR client on every connect.
|
||||
func TestRepeatedReconnectConsumesOnlyThatClientsBudget(t *testing.T) {
|
||||
a := newProxiedAPI(t)
|
||||
|
||||
const reconnecting = "203.0.113.60"
|
||||
a.exhaust(t, proxyAddr, reconnecting, "reconnector")
|
||||
|
||||
// Nine other clients are untouched by it.
|
||||
for i := 0; i < 9; i++ {
|
||||
other := fmt.Sprintf("198.51.100.%d", i+1)
|
||||
if code := a.register(t, proxyAddr, other, fmt.Sprintf("other-%d", i)); code != http.StatusCreated {
|
||||
t.Fatalf("client %s: status %d, want 201", other, code)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Spoofing still buys nothing ────────────────────────────────────────── */
|
||||
|
||||
// A caller reaching the API directly — not through the proxy — cannot escape
|
||||
// its bucket by varying X-Forwarded-For. It gets ONE budget however many
|
||||
// different values it sends.
|
||||
func TestUntrustedCallerCannotEscapeItsBucketByForging(t *testing.T) {
|
||||
a := newProxiedAPI(t)
|
||||
|
||||
const direct = "198.51.100.200:40000" // outside proxyNetwork
|
||||
|
||||
allowed := 0
|
||||
for i := 0; i < ratelimit.OAuthRegister.Limit*2; i++ {
|
||||
// A different forged client on every single request.
|
||||
code := a.register(t, direct, fmt.Sprintf("203.0.113.%d", i+100), fmt.Sprintf("forger-%d", i))
|
||||
if code == http.StatusTooManyRequests {
|
||||
break
|
||||
}
|
||||
if code != http.StatusCreated {
|
||||
t.Fatalf("attempt %d: unexpected status %d", i, code)
|
||||
}
|
||||
allowed++
|
||||
}
|
||||
|
||||
if allowed != ratelimit.OAuthRegister.Limit {
|
||||
t.Errorf("a forging caller got %d registrations, want %d — the header bought extra budget",
|
||||
allowed, ratelimit.OAuthRegister.Limit)
|
||||
}
|
||||
|
||||
var buckets int
|
||||
if err := a.h.Pool.QueryRow(t.Context(),
|
||||
`SELECT count(*) FROM rate_limits WHERE bucket LIKE 'oauth.register:%'`).Scan(&buckets); err != nil {
|
||||
t.Fatalf("count buckets: %v", err)
|
||||
}
|
||||
if buckets != 1 {
|
||||
t.Errorf("a forging caller produced %d buckets, want exactly 1", buckets)
|
||||
}
|
||||
}
|
||||
|
||||
// Claiming to be the trusted proxy does not make a caller trusted.
|
||||
func TestClaimingToBeTheProxyDoesNotWork(t *testing.T) {
|
||||
a := newProxiedAPI(t)
|
||||
|
||||
const direct = "198.51.100.201:40000"
|
||||
// The forged chain ends in the proxy's own address, which is what an
|
||||
// attacker who has read this file would try.
|
||||
if code := a.register(t, direct, "203.0.113.9, 10.42.0.1", "impostor"); code != http.StatusCreated {
|
||||
t.Fatalf("setup: status %d", code)
|
||||
}
|
||||
|
||||
var bucket string
|
||||
if err := a.h.Pool.QueryRow(t.Context(),
|
||||
`SELECT bucket FROM rate_limits WHERE bucket LIKE 'oauth.register:%'`).Scan(&bucket); err != nil {
|
||||
t.Fatalf("read bucket: %v", err)
|
||||
}
|
||||
// The bucket must be the DIRECT caller's address, not the forged one.
|
||||
want := "oauth.register:" + ratelimit.Subject("198.51.100.201")
|
||||
if bucket != want {
|
||||
t.Errorf("bucket = %q, want %q — a forged chain was believed", bucket, want)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── The default posture is unchanged ───────────────────────────────────── */
|
||||
|
||||
// With no trusted proxies configured — the default, and how every other test in
|
||||
// this package runs — the forwarded header is ignored and callers share the
|
||||
// peer's bucket exactly as before.
|
||||
func TestWithoutTrustedProxiesCallersShareThePeerBucket(t *testing.T) {
|
||||
a := newOAuthAPI(t) // no TrustedProxies in its config
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
body := fmt.Sprintf(
|
||||
`{"client_name":"unproxied-%d","redirect_uris":["https://client.example.test/cb"]}`, i)
|
||||
req, err := http.NewRequest("POST", "/oauth/register", strings.NewReader(body))
|
||||
if err != nil {
|
||||
t.Fatalf("build request: %v", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.RemoteAddr = "192.0.2.10:5000"
|
||||
req.Header.Set("X-Forwarded-For", fmt.Sprintf("203.0.113.%d", i+1))
|
||||
rec := httptest.NewRecorder()
|
||||
a.handler.ServeHTTP(rec, req)
|
||||
if rec.Code != http.StatusCreated {
|
||||
t.Fatalf("attempt %d: status %d", i, rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
var buckets int
|
||||
if err := a.h.Pool.QueryRow(t.Context(),
|
||||
`SELECT count(*) FROM rate_limits WHERE bucket LIKE 'oauth.register:%'`).Scan(&buckets); err != nil {
|
||||
t.Fatalf("count buckets: %v", err)
|
||||
}
|
||||
if buckets != 1 {
|
||||
t.Errorf("%d buckets with no trusted proxy configured, want 1 — the header was read", buckets)
|
||||
}
|
||||
}
|
||||
@@ -1,9 +1,6 @@
|
||||
package httpserver
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
@@ -23,12 +20,13 @@ import (
|
||||
// one instance needs shared state — Redis, or the database — and this
|
||||
// package is the seam where that goes: attemptLimiter is an implementation
|
||||
// detail behind Allow/Fail/Reset.
|
||||
// - It trusts net/http's RemoteAddr for the client address. Behind a reverse
|
||||
// proxy every request appears to come from the proxy, so the per-address
|
||||
// budget becomes global. Reading X-Forwarded-For instead would be worse,
|
||||
// not better, until there is a trusted-proxy list to validate it against —
|
||||
// a client can send that header itself and mint a fresh budget per request.
|
||||
// Deploying behind a proxy means adding that list first.
|
||||
// - The per-address budget is only as good as the address. That used to be
|
||||
// net/http's RemoteAddr, which behind a reverse proxy is the proxy on
|
||||
// every request and makes this budget global. It is now resolved by
|
||||
// proxyTrust.clientAddr (clientip.go), which reads a forwarded address
|
||||
// when — and only when — the immediate peer is a configured trusted
|
||||
// proxy. An unconfigured deployment still gets RemoteAddr, so a proxied
|
||||
// deployment must set HTTP_TRUSTED_PROXIES for this limit to be per-user.
|
||||
// - It is memory-bounded by pruning, not by a hard cap, so a flood from many
|
||||
// distinct addresses grows the map until the next prune.
|
||||
//
|
||||
@@ -139,16 +137,3 @@ func (l *attemptLimiter) pruneLocked(now time.Time) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// clientAddr is the key for per-address limiting.
|
||||
//
|
||||
// The port is stripped: a browser uses a new source port for every connection,
|
||||
// so keying on host:port would give each attempt its own budget and limit
|
||||
// nothing at all.
|
||||
func clientAddr(r *http.Request) string {
|
||||
host, _, err := net.SplitHostPort(strings.TrimSpace(r.RemoteAddr))
|
||||
if err != nil {
|
||||
return strings.TrimSpace(r.RemoteAddr)
|
||||
}
|
||||
return host
|
||||
}
|
||||
|
||||
@@ -37,6 +37,7 @@ import (
|
||||
"github.com/krow/krow-backend/go-api/internal/db"
|
||||
"github.com/krow/krow-backend/go-api/internal/definition"
|
||||
"github.com/krow/krow-backend/go-api/internal/knowledge"
|
||||
"github.com/krow/krow-backend/go-api/internal/ratelimit"
|
||||
"github.com/krow/krow-backend/go-api/internal/runtime"
|
||||
"github.com/krow/krow-backend/go-api/internal/service"
|
||||
"github.com/krow/krow-backend/go-api/internal/tools"
|
||||
@@ -83,7 +84,18 @@ type Server struct {
|
||||
users auth.UserStore
|
||||
credentials *auth.Credentials
|
||||
loginByEmail *attemptLimiter
|
||||
loginByAddr *attemptLimiter
|
||||
|
||||
// limiter bounds the OAuth and MCP routes, shared across instances via
|
||||
// Postgres. Nil when those routes are not registered — see mcplimit.go,
|
||||
// where a nil limiter means the middleware is not installed at all rather
|
||||
// than installed and permissive.
|
||||
limiter *ratelimit.Limiter
|
||||
loginByAddr *attemptLimiter
|
||||
|
||||
// trust resolves a request to the address its per-address limits are keyed
|
||||
// by, reading a forwarded address only from a configured proxy. Wired once
|
||||
// here so no handler can be given a different notion of who called.
|
||||
trust proxyTrust
|
||||
|
||||
// now is injectable so tests can drive expiry without sleeping.
|
||||
now func() time.Time
|
||||
@@ -223,6 +235,7 @@ func New(cfg *config.Config, database *db.DB, log *slog.Logger, opts ...Option)
|
||||
credentials: auth.NewCredentials(users),
|
||||
loginByEmail: newAttemptLimiter(o.perEmail, o.loginWindow, o.now),
|
||||
loginByAddr: newAttemptLimiter(o.perAddress, o.loginWindow, o.now),
|
||||
trust: newProxyTrust(cfg.HTTP.TrustedProxies),
|
||||
now: o.now,
|
||||
}
|
||||
|
||||
@@ -277,11 +290,23 @@ func New(cfg *config.Config, database *db.DB, log *slog.Logger, opts ...Option)
|
||||
}
|
||||
s.definitions = s.definitions.WithCuratedAgents(curated)
|
||||
|
||||
// The shared limiter, built only when the routes that use it exist. The
|
||||
// existing in-process login limiter is untouched: it guards a different
|
||||
// thing (failed password attempts) with a different model (count failures,
|
||||
// reset on success), and replacing it is not this change's business.
|
||||
if cfg.OAuth.Enabled() {
|
||||
s.limiter = ratelimit.New(database.Pool)
|
||||
}
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("GET /health", s.handleHealth)
|
||||
s.endpoints = s.routeAuth(mux) + s.routeResources(mux) + s.routeMe(mux) +
|
||||
s.routeDefinitions(mux) + s.routeWorkflows(mux) + s.routeOwliver(mux) +
|
||||
s.routeRuns(mux) + s.routeVersion(mux) + s.routeTools(mux)
|
||||
s.routeRuns(mux) + s.routeVersion(mux) + s.routeTools(mux) +
|
||||
// The MCP surface and the OAuth server behind it. Both return 0 and
|
||||
// register nothing when OAUTH_ISSUER and MCP_RESOURCE are unset, which
|
||||
// is every deployment that has not asked for them.
|
||||
s.routeOAuth(mux) + s.routeMCP(mux)
|
||||
|
||||
handler := jsonErrors(mux)
|
||||
// Authentication sits where devOrgMiddleware used to, so every route below
|
||||
|
||||
Reference in New Issue
Block a user