package auth import ( "context" "net/http" "net/http/httptest" "testing" ) // errResolver rejects every request, standing in for a real resolver that finds // no valid session. type errResolver struct{ err error } func (e errResolver) Resolve(*http.Request) (string, error) { return "", e.err } func TestUserIDRoundTrip(t *testing.T) { ctx := WithUser(context.Background(), "alice") if got := UserID(ctx); got != "alice" { t.Fatalf("UserID = %q, want alice", got) } } // A request that never passed through the middleware must report no user rather // than panicking — every query is `WHERE user_id = ?`, so an empty id fails // closed (matches nothing) instead of falling back to some default account. func TestUserIDAbsentIsEmpty(t *testing.T) { if got := UserID(context.Background()); got != "" { t.Fatalf("UserID on bare context = %q, want empty", got) } } func TestMiddlewareInjectsResolvedUser(t *testing.T) { var seen string h := Middleware(StaticResolver("local"))(http.HandlerFunc( func(_ http.ResponseWriter, r *http.Request) { seen = UserID(r.Context()) }, )) rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil)) if seen != "local" { t.Fatalf("handler saw user %q, want local", seen) } if rec.Code != http.StatusOK { t.Fatalf("status = %d, want 200", rec.Code) } } // Both rejection paths — an explicit error and a silent empty id — must 401 // without ever entering the handler. The empty case matters most: a resolver // that returns ("", nil) by mistake would otherwise hand handlers an empty user // id, and while that fails closed at the SQL layer, it should never get there. func TestMiddlewareRejectsUnresolved(t *testing.T) { for name, res := range map[string]Resolver{ "resolver error": errResolver{err: http.ErrNoCookie}, "empty user id": StaticResolver(""), } { t.Run(name, func(t *testing.T) { called := false h := Middleware(res)(http.HandlerFunc( func(http.ResponseWriter, *http.Request) { called = true }, )) rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil)) if called { t.Fatal("handler ran for an unauthenticated request") } if rec.Code != http.StatusUnauthorized { t.Fatalf("status = %d, want 401", rec.Code) } }) } }