~bigbes/core-go

bfc744fce7546de301f8a2d2cd822c42fa166f0e — Drew DeVault 5 years ago e164f26
Add basic auth tests
7 files changed, 256 insertions(+), 6 deletions(-)

M auth/middleware.go
A auth/middleware_test.go
M auth/token.go
A auth/token_test.go
M database/middleware.go
M go.mod
M go.sum
M auth/middleware.go => auth/middleware.go +1 -2
@@ 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 <token>'`, http.StatusForbidden)
				authError(w, `Authorization header is required. Expected 'Authorization: Bearer <token>'`, 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
			}

		})
	}
}

A auth/middleware_test.go => auth/middleware_test.go +144 -0
@@ 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
}

M auth/token.go => auth/token.go +2 -2
@@ 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 {

A auth/token_test.go => auth/token_test.go +101 -0
@@ 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)
}

M database/middleware.go => database/middleware.go +5 -2
@@ 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 {

M go.mod => go.mod +1 -0
@@ 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

M go.sum => go.sum +2 -0
@@ 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=