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=