From bfc744fce7546de301f8a2d2cd822c42fa166f0e Mon Sep 17 00:00:00 2001 From: Drew DeVault Date: Tue, 6 Oct 2020 13:02:51 -0400 Subject: [PATCH] Add basic auth tests --- auth/middleware.go | 3 +- auth/middleware_test.go | 144 ++++++++++++++++++++++++++++++++++++++++ auth/token.go | 4 +- auth/token_test.go | 101 ++++++++++++++++++++++++++++ database/middleware.go | 7 +- go.mod | 1 + go.sum | 2 + 7 files changed, 256 insertions(+), 6 deletions(-) create mode 100644 auth/middleware_test.go create mode 100644 auth/token_test.go diff --git a/auth/middleware.go b/auth/middleware.go index 728b8c031bdce1082c228c558be12cc29584f855..54a710183b1c66df7ed8b86d1e958ddcd354612f 100644 --- a/auth/middleware.go +++ b/auth/middleware.go @@ -686,7 +686,7 @@ func Middleware(conf ini.File, apiconf string) func(http.Handler) http.Handler { auth := r.Header.Get("Authorization") if auth == "" { - authError(w, `Authorization header is required. Expected 'Authorization: Bearer '`, http.StatusForbidden) + authError(w, `Authorization header is required. Expected 'Authorization: Bearer '`, http.StatusUnauthorized) return } @@ -722,7 +722,6 @@ func Middleware(conf ini.File, apiconf string) func(http.Handler) http.Handler { authError(w, "Invalid Authorization header", http.StatusBadRequest) return } - }) } } diff --git a/auth/middleware_test.go b/auth/middleware_test.go new file mode 100644 index 0000000000000000000000000000000000000000..2a61e6723a3f00577956244da75ab719f103cfb1 --- /dev/null +++ b/auth/middleware_test.go @@ -0,0 +1,144 @@ +package auth + +import ( + "context" + "encoding/json" + "net/http" + "strings" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/stretchr/testify/assert" + "github.com/vaughan0/go-ini" + + "git.sr.ht/~sircmpwn/core-go/crypto" + "git.sr.ht/~sircmpwn/core-go/database" +) + +func TestNoAuthorization(t *testing.T) { + mw, _, next := middleware() + req, err := http.NewRequestWithContext(context.Background(), "POST", + "https://example.org/query", + strings.NewReader(`{"query": "query { me { id } }"}`)) + assert.Nil(t, err) + req.Header.Add("Content-Type", "application/json") + resp := &TestResponse{T: t} + mw(resp, req) + assert.False(t, *next) + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) +} + +func TestCookie(t *testing.T) { + mw, subctx, next := middleware() + ctx, mock := dbctx() + req, err := http.NewRequestWithContext(ctx, "POST", + "https://example.org/query", + strings.NewReader(`{"query": "query { me { id } }"}`)) + assert.Nil(t, err) + req.Header.Add("Content-Type", "application/json") + + cookie := AuthCookie{ + Name: "jdoe", + } + payload, err := json.Marshal(&cookie) + assert.Nil(t, err) + + req.AddCookie(&http.Cookie{ + Name: "sr.ht.unified-login.v1", + Value: string(crypto.Encrypt(payload)), + }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`). + WithArgs("jdoe"). + WillReturnRows(sqlmock.NewRows([]string{ + "id", "username", "created", "updated", "email", "user_type", + "url", "location", "bio", "suspension_notice", + }). + AddRow(1337, "jdoe", time.Now().UTC(), time.Now().UTC(), + "jdoe@example.org", "active_paying", + "https://example.org", nil, nil, nil)) + mock.ExpectCommit() + + resp := &TestResponse{T: t} + mw(resp, req) + assert.True(t, *next) + assert.Nil(t, mock.ExpectationsWereMet()) + + auth := ForContext(*subctx) + assert.Equal(t, auth.AuthMethod, AUTH_COOKIE) + assert.Equal(t, auth.UserID, 1337) + assert.Equal(t, auth.Username, "jdoe") + assert.Equal(t, auth.Email, "jdoe@example.org") +} + +var conf ini.File + +func init() { + var err error + conf, err = ini.Load(strings.NewReader(` +[webhooks] +private-key=ebzsjPaN6E13ln/FeNWly1C92q6bVMVdOnDo1HPl5fc= + +[sr.ht] +network-key=tbuG-7Vh44vrDq1L_HKWkHnWrDOtJhEkPKPiauaLeuk= + +[test::api] +internal-ipnet=127.0.0.1/24,::1/64`)) + if err != nil { + panic(err) + } + crypto.InitCrypto(conf) +} + +func middleware() (http.HandlerFunc, *context.Context, *bool) { + called := false + var ctx context.Context + next := (func(called *bool, ctx *context.Context) http.HandlerFunc { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + *called = true + *ctx = r.Context() + }) + })(&called, &ctx) + return Middleware(conf, "test::api")(next).ServeHTTP, &ctx, &called +} + +func dbctx() (context.Context, sqlmock.Sqlmock) { + db, mock, err := sqlmock.New() + if err != nil { + panic(err) + } + ctx := database.Context(context.Background(), db) + return ctx, mock +} + +type TestResponse struct { + T *testing.T + + RespHeader http.Header + Payload []byte + StatusCode int +} + +func (tr *TestResponse) Header() http.Header { + if tr.RespHeader == nil { + tr.RespHeader = make(http.Header) + } + return tr.RespHeader +} + +func (tr *TestResponse) Write(payload []byte) (int, error) { + if tr.StatusCode == 0 { + tr.WriteHeader(http.StatusOK) + } + + assert.Nil(tr.T, tr.Payload) + tr.Payload = payload + return len(payload), nil +} + +func (tr *TestResponse) WriteHeader(statusCode int) { + assert.Zero(tr.T, tr.StatusCode) + tr.StatusCode = statusCode +} diff --git a/auth/token.go b/auth/token.go index 1e21e878bb85a129ff30017ecb34adb7343671f7..ca2baeacc15827d9ba32fc4050c07cdfa6f7ca5a 100644 --- a/auth/token.go +++ b/auth/token.go @@ -16,11 +16,11 @@ const TokenVersion uint = 0 type Timestamp int64 func (t Timestamp) Time() time.Time { - return time.Unix(int64(t), 0) + return time.Unix(int64(t), 0).UTC() } func ToTimestamp(t time.Time) Timestamp { - return Timestamp(t.Unix()) + return Timestamp(t.UTC().Unix()) } type OAuth2Token struct { diff --git a/auth/token_test.go b/auth/token_test.go new file mode 100644 index 0000000000000000000000000000000000000000..e33ae096229bed177c6fea74f15258c0e01d0055 --- /dev/null +++ b/auth/token_test.go @@ -0,0 +1,101 @@ +package auth + +import ( + "encoding/base64" + "strings" + "testing" + "time" + + "git.sr.ht/~sircmpwn/go-bare" + "github.com/stretchr/testify/assert" + "github.com/vaughan0/go-ini" + + "git.sr.ht/~sircmpwn/core-go/crypto" +) + +func init() { + config, err := ini.Load(strings.NewReader(` +[webhooks] +private-key=ebzsjPaN6E13ln/FeNWly1C92q6bVMVdOnDo1HPl5fc= + +[sr.ht] +network-key=tbuG-7Vh44vrDq1L_HKWkHnWrDOtJhEkPKPiauaLeuk=`)) + if err != nil { + panic(err) + } + crypto.InitCrypto(config) +} + +func TestEncode(t *testing.T) { + ot := &OAuth2Token{ + Version: TokenVersion, + Expires: ToTimestamp(time.Now().Add(30 * time.Minute)), + Grants: "", + ClientID: "", + Username: "jdoe", + } + token := ot.Encode() + bytes, err := base64.RawStdEncoding.DecodeString(token) + assert.Nil(t, err) + + mac := bytes[len(bytes)-32:] + payload := bytes[:len(bytes)-32] + assert.True(t, crypto.HMACVerify(payload, mac)) + + var ot2 OAuth2Token + err = bare.Unmarshal(payload, &ot2) + assert.Nil(t, err) + assert.Equal(t, ot.Version, ot2.Version) + assert.Equal(t, ot.Expires, ot2.Expires) + assert.Equal(t, ot.Grants, ot2.Grants) + assert.Equal(t, ot.ClientID, ot2.ClientID) + assert.Equal(t, ot.Username, ot2.Username) +} + +func TestDecode(t *testing.T) { + ot := &OAuth2Token{ + Version: TokenVersion, + Expires: ToTimestamp(time.Now().Add(30 * time.Minute)), + Grants: "", + ClientID: "", + Username: "jdoe", + } + token := ot.Encode() + ot2 := DecodeToken(token) + assert.NotNil(t, ot2) + assert.Equal(t, ot.Version, ot2.Version) + assert.Equal(t, ot.Expires, ot2.Expires) + assert.Equal(t, ot.Grants, ot2.Grants) + assert.Equal(t, ot.ClientID, ot2.ClientID) + assert.Equal(t, ot.Username, ot2.Username) + + // Expired token: + ot = &OAuth2Token{ + Version: TokenVersion, + Expires: ToTimestamp(time.Now().Add(-30 * time.Minute)), + Grants: "", + ClientID: "", + Username: "jdoe", + } + token = ot.Encode() + ot2 = DecodeToken(token) + assert.Nil(t, ot2) + + // Invalid MAC: + ot = &OAuth2Token{ + Version: TokenVersion, + Expires: ToTimestamp(time.Now().Add(30 * time.Minute)), + Grants: "", + ClientID: "", + Username: "jdoe", + } + plain, err := bare.Marshal(ot) + assert.Nil(t, err) + mac := crypto.HMAC(plain) + ot.Username = "rdoe" + plain, err = bare.Marshal(ot) + assert.Nil(t, err) + token = base64.RawStdEncoding.EncodeToString(append(plain, mac...)) + ot2 = DecodeToken(token) + assert.Nil(t, ot2) +} diff --git a/database/middleware.go b/database/middleware.go index 9e2761e48225b37b3962d61a03ea174ebd9f65c3..ee9de28887fd0539e299b6bf2113a853eb0a0f5e 100644 --- a/database/middleware.go +++ b/database/middleware.go @@ -16,14 +16,17 @@ type contextKey struct { func Middleware(db *sql.DB) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - ctx := context.WithValue(r.Context(), dbCtxKey, db) - + ctx := Context(r.Context(), db) r = r.WithContext(ctx) next.ServeHTTP(w, r) }) } } +func Context(ctx context.Context, db *sql.DB) context.Context { + return context.WithValue(ctx, dbCtxKey, db) +} + func ForContext(ctx context.Context) (*sql.Conn, error) { raw, ok := ctx.Value(dbCtxKey).(*sql.DB) if !ok { diff --git a/go.mod b/go.mod index 2b36d8db0dce7f03838b305a1a978f93e07f693f..d733cc3e0740db5921e76839f6aff2a9889ed86c 100644 --- a/go.mod +++ b/go.mod @@ -7,6 +7,7 @@ require ( git.sr.ht/~sircmpwn/go-bare v0.0.0-20200812160916-d2c72e1a5018 git.sr.ht/~sircmpwn/gql.sr.ht v0.0.0-20200924184215-6a1a8f1031f3 github.com/99designs/gqlgen v0.13.0 + github.com/DATA-DOG/go-sqlmock v1.5.0 github.com/Masterminds/squirrel v1.4.0 github.com/fernet/fernet-go v0.0.0-20191111064656-eff2850e6001 github.com/go-chi/chi v4.1.2+incompatible diff --git a/go.sum b/go.sum index f065687c2940bb30ddd5acd5dd5b5fc0503dd3b6..9851eddf8cf708bba835df96e4671a8452ae0f68 100644 --- a/go.sum +++ b/go.sum @@ -8,6 +8,8 @@ github.com/99designs/gqlgen v0.11.4-0.20200512031635-40570d1b4d70/go.mod h1:RgX5 github.com/99designs/gqlgen v0.13.0 h1:haLTcUp3Vwp80xMVEg5KRNwzfUrgFdRmtBY8fuB8scA= github.com/99designs/gqlgen v0.13.0/go.mod h1:NV130r6f4tpRWuAI+zsrSdooO/eWUv+Gyyoi3rEfXIk= github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= +github.com/DATA-DOG/go-sqlmock v1.5.0 h1:Shsta01QNfFxHCfpW6YH2STWB0MudeXXEWMr20OEh60= +github.com/DATA-DOG/go-sqlmock v1.5.0/go.mod h1:f/Ixk793poVmq4qj/V1dPUg2JEAKC73Q5eFN3EC/SaM= github.com/Masterminds/squirrel v1.4.0 h1:he5i/EXixZxrBUWcxzDYMiju9WZ3ld/l7QBNuo/eN3w= github.com/Masterminds/squirrel v1.4.0/go.mod h1:yaPeOnPG5ZRwL9oKdTsO/prlkPbXWZlRVMQ/gGlzIuA= github.com/agnivade/levenshtein v1.0.1/go.mod h1:CURSv5d9Uaml+FovSIICkLbAUZ9S4RqaHDIsdSBg7lM=