Codex first iteration
This commit is contained in:
@@ -0,0 +1,174 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Role string
|
||||
|
||||
const (
|
||||
Student Role = "student"
|
||||
Instructor Role = "instructor"
|
||||
CookieName = "preface_session"
|
||||
)
|
||||
|
||||
type Claims struct {
|
||||
Role Role `json:"role"`
|
||||
SessionID string `json:"session_id,omitempty"`
|
||||
IssuedAt int64 `json:"issued_at"`
|
||||
ExpiresAt int64 `json:"expires_at"`
|
||||
Version string `json:"credential_version"`
|
||||
CSRFToken string `json:"csrf_token"`
|
||||
}
|
||||
type Service struct {
|
||||
studentPIN, instructorPIN string
|
||||
secret []byte
|
||||
duration time.Duration
|
||||
secure bool
|
||||
limiter *Limiter
|
||||
}
|
||||
|
||||
func New(studentPIN, instructorPIN, secret string, duration time.Duration, secure bool) *Service {
|
||||
return &Service{studentPIN: studentPIN, instructorPIN: instructorPIN, secret: []byte(secret), duration: duration, secure: secure, limiter: NewLimiter(5, 5*time.Minute)}
|
||||
}
|
||||
func constantEqual(a, b string) bool { return hmac.Equal([]byte(a), []byte(b)) }
|
||||
func (s *Service) VerifyInstructorPIN(pin string) bool { return constantEqual(pin, s.instructorPIN) }
|
||||
func randomToken() string {
|
||||
b := make([]byte, 24)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return base64.RawURLEncoding.EncodeToString(b)
|
||||
}
|
||||
func (s *Service) version(role Role) string {
|
||||
pin := s.studentPIN
|
||||
if role == Instructor {
|
||||
pin = s.instructorPIN
|
||||
}
|
||||
m := hmac.New(sha256.New, s.secret)
|
||||
m.Write([]byte(string(role) + pin))
|
||||
return base64.RawURLEncoding.EncodeToString(m.Sum(nil)[:16])
|
||||
}
|
||||
func (s *Service) Authenticate(role Role, pin, key string) (Claims, error) {
|
||||
if !s.limiter.Allow(key) {
|
||||
return Claims{}, errors.New("too many attempts")
|
||||
}
|
||||
expected := s.studentPIN
|
||||
if role == Instructor {
|
||||
expected = s.instructorPIN
|
||||
}
|
||||
if !constantEqual(pin, expected) {
|
||||
s.limiter.Fail(key)
|
||||
return Claims{}, errors.New("invalid PIN")
|
||||
}
|
||||
s.limiter.Success(key)
|
||||
now := time.Now()
|
||||
c := Claims{Role: role, IssuedAt: now.Unix(), ExpiresAt: now.Add(s.duration).Unix(), Version: s.version(role), CSRFToken: randomToken()}
|
||||
if role == Student {
|
||||
c.SessionID = randomToken()
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
func (s *Service) Sign(c Claims) string {
|
||||
b, _ := json.Marshal(c)
|
||||
p := base64.RawURLEncoding.EncodeToString(b)
|
||||
m := hmac.New(sha256.New, s.secret)
|
||||
m.Write([]byte(p))
|
||||
return p + "." + base64.RawURLEncoding.EncodeToString(m.Sum(nil))
|
||||
}
|
||||
func (s *Service) Parse(raw string) (Claims, error) {
|
||||
var c Claims
|
||||
parts := strings.Split(raw, ".")
|
||||
if len(parts) != 2 {
|
||||
return c, errors.New("bad session")
|
||||
}
|
||||
sig, err := base64.RawURLEncoding.DecodeString(parts[1])
|
||||
if err != nil {
|
||||
return c, err
|
||||
}
|
||||
m := hmac.New(sha256.New, s.secret)
|
||||
m.Write([]byte(parts[0]))
|
||||
if !hmac.Equal(sig, m.Sum(nil)) {
|
||||
return c, errors.New("bad signature")
|
||||
}
|
||||
b, err := base64.RawURLEncoding.DecodeString(parts[0])
|
||||
if err != nil {
|
||||
return c, err
|
||||
}
|
||||
if err = json.Unmarshal(b, &c); err != nil {
|
||||
return c, err
|
||||
}
|
||||
if time.Now().Unix() >= c.ExpiresAt || c.Version != s.version(c.Role) {
|
||||
return c, errors.New("expired session")
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
func (s *Service) SetCookie(w http.ResponseWriter, c Claims) {
|
||||
http.SetCookie(w, &http.Cookie{Name: CookieName, Value: s.Sign(c), Path: "/", HttpOnly: true, Secure: s.secure, SameSite: http.SameSiteLaxMode, MaxAge: int(s.duration.Seconds())})
|
||||
}
|
||||
func (s *Service) ClearCookie(w http.ResponseWriter) {
|
||||
http.SetCookie(w, &http.Cookie{Name: CookieName, Path: "/", HttpOnly: true, Secure: s.secure, SameSite: http.SameSiteLaxMode, MaxAge: -1})
|
||||
}
|
||||
func (s *Service) FromRequest(r *http.Request) (Claims, error) {
|
||||
c, err := r.Cookie(CookieName)
|
||||
if err != nil {
|
||||
return Claims{}, err
|
||||
}
|
||||
return s.Parse(c.Value)
|
||||
}
|
||||
|
||||
type contextKey struct{}
|
||||
|
||||
func WithClaims(r *http.Request, c Claims) *http.Request {
|
||||
return r.WithContext(context.WithValue(r.Context(), contextKey{}, c))
|
||||
}
|
||||
func ClaimsFrom(r *http.Request) (Claims, bool) {
|
||||
c, ok := r.Context().Value(contextKey{}).(Claims)
|
||||
return c, ok
|
||||
}
|
||||
|
||||
type attempt struct {
|
||||
failures int
|
||||
reset time.Time
|
||||
}
|
||||
type Limiter struct {
|
||||
mu sync.Mutex
|
||||
max int
|
||||
window time.Duration
|
||||
entries map[string]attempt
|
||||
}
|
||||
|
||||
func NewLimiter(max int, window time.Duration) *Limiter {
|
||||
return &Limiter{max: max, window: window, entries: map[string]attempt{}}
|
||||
}
|
||||
func (l *Limiter) Allow(k string) bool {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
a := l.entries[k]
|
||||
if time.Now().After(a.reset) {
|
||||
delete(l.entries, k)
|
||||
return true
|
||||
}
|
||||
return a.failures < l.max
|
||||
}
|
||||
func (l *Limiter) Fail(k string) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
a := l.entries[k]
|
||||
if time.Now().After(a.reset) {
|
||||
a = attempt{reset: time.Now().Add(l.window)}
|
||||
}
|
||||
a.failures++
|
||||
l.entries[k] = a
|
||||
}
|
||||
func (l *Limiter) Success(k string) { l.mu.Lock(); defer l.mu.Unlock(); delete(l.entries, k) }
|
||||
@@ -0,0 +1,37 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestSessionsAndRotation(t *testing.T) {
|
||||
s := New("1234", "9876", string(make([]byte, 32)), time.Hour, false)
|
||||
c, err := s.Authenticate(Student, "1234", "ip")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if c.SessionID == "" || c.CSRFToken == "" {
|
||||
t.Fatal("missing random claims")
|
||||
}
|
||||
raw := s.Sign(c)
|
||||
got, err := s.Parse(raw)
|
||||
if err != nil || got.Role != Student {
|
||||
t.Fatalf("parse: %#v %v", got, err)
|
||||
}
|
||||
rotated := New("4321", "9876", string(make([]byte, 32)), time.Hour, false)
|
||||
if _, err := rotated.Parse(raw); err == nil {
|
||||
t.Fatal("PIN rotation did not invalidate session")
|
||||
}
|
||||
}
|
||||
func TestWrongPINAndExpiry(t *testing.T) {
|
||||
s := New("1234", "9876", string(make([]byte, 32)), time.Hour, false)
|
||||
if _, err := s.Authenticate(Instructor, "bad", "ip"); err == nil {
|
||||
t.Fatal("wrong PIN accepted")
|
||||
}
|
||||
c, _ := s.Authenticate(Student, "1234", "another")
|
||||
c.ExpiresAt = time.Now().Add(-time.Second).Unix()
|
||||
if _, err := s.Parse(s.Sign(c)); err == nil {
|
||||
t.Fatal("expired session accepted")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user