Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions pkg/kratos/cookies.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
// Copyright 2024 Canonical Ltd.
// SPDX-License-Identifier: AGPL-3.0

// Package kratos provides state cookie management for login flow state tracking.
package kratos

import (
Expand Down Expand Up @@ -34,6 +35,7 @@ type FlowStateCookie struct {
BackupCodeUsed bool `json:"bc,omitempty"`
}

// SetStateCookie sets a state cookie on the HTTP response.
func (a *AuthCookieManager) SetStateCookie(w http.ResponseWriter, state FlowStateCookie) error {
rawState, err := json.Marshal(state)
if err != nil {
Expand All @@ -42,6 +44,7 @@ func (a *AuthCookieManager) SetStateCookie(w http.ResponseWriter, state FlowStat
return a.setCookie(w, stateCookieName, string(rawState), defaultCookiePath, a.cookieTTL, http.SameSiteLaxMode)
}

// GetStateCookie retrieves a state cookie from the HTTP request.
func (a *AuthCookieManager) GetStateCookie(r *http.Request) (FlowStateCookie, error) {
var ret FlowStateCookie
c, err := a.getCookie(r, stateCookieName)
Expand All @@ -52,10 +55,12 @@ func (a *AuthCookieManager) GetStateCookie(r *http.Request) (FlowStateCookie, er
return ret, err
}

// ClearStateCookie clears the state cookie from the HTTP response.
func (a *AuthCookieManager) ClearStateCookie(w http.ResponseWriter) {
a.clearCookie(w, stateCookieName, defaultCookiePath)
}

// setCookie sets a cookie on the HTTP response with the specified parameters.
func (a *AuthCookieManager) setCookie(w http.ResponseWriter, name, value string, path string, ttl time.Duration, sameSitePolicy http.SameSite) error {
if value == "" {
return nil
Expand Down Expand Up @@ -83,6 +88,7 @@ func (a *AuthCookieManager) setCookie(w http.ResponseWriter, name, value string,
return nil
}

// clearCookie removes a cookie from the HTTP response.
func (a *AuthCookieManager) clearCookie(w http.ResponseWriter, name string, path string) {
http.SetCookie(w, &http.Cookie{
Name: name,
Expand All @@ -96,6 +102,7 @@ func (a *AuthCookieManager) clearCookie(w http.ResponseWriter, name string, path
})
}

// getCookie retrieves a cookie value from the HTTP request.
func (a *AuthCookieManager) getCookie(r *http.Request, name string) (string, error) {
cookie, err := r.Cookie(name)
if err != nil {
Expand All @@ -111,6 +118,7 @@ func (a *AuthCookieManager) getCookie(r *http.Request, name string) (string, err
return value, nil
}

// NewAuthCookieManager creates a new AuthCookieManager instance.
func NewAuthCookieManager(
cookieTTLSeconds int,
encrypt EncryptInterface,
Expand Down
263 changes: 148 additions & 115 deletions pkg/kratos/cookies_test.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
// Copyright 2024 Canonical Ltd.
// SPDX-License-Identifier: AGPL-3.0

// Package kratos provides unit tests for cookie management functionality.
package kratos

import (
Expand Down Expand Up @@ -29,147 +30,179 @@ func findCookie(name string, cookies []*http.Cookie) (*http.Cookie, bool) {
}

func TestAuthCookieManager_ClearStateCookie(t *testing.T) {
ctrl := gomock.NewController(t)

mockLogger := NewMockLoggerInterface(ctrl)
mockEncrypt := NewMockEncryptInterface(ctrl)

mockRequest := httptest.NewRequest(http.MethodGet, "/", nil)
mockRequest.AddCookie(&http.Cookie{Name: "state"})

mockResponse := httptest.NewRecorder()

manager := NewAuthCookieManager(5, mockEncrypt, mockLogger)
manager.ClearStateCookie(mockResponse)

c, _ := findCookie("login_ui_state", mockResponse.Result().Cookies())

if c.Expires != epoch {
t.Fatal("did not clear state cookie")
tests := []struct {
name string
}{
{
name: "ClearState",
},
}
Comment on lines +33 to 39

Copilot AI Feb 2, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The table-driven test structure for TestAuthCookieManager_ClearStateCookie contains only a single test case with no variations. This adds unnecessary complexity without benefit. Either add multiple test cases to justify the table-driven approach, or revert to a simple test function.

Copilot uses AI. Check for mistakes.
}

func TestAuthCookieManager_GetStateCookie(t *testing.T) {
ctrl := gomock.NewController(t)

mockLogger := NewMockLoggerInterface(ctrl)
mockEncrypt := NewMockEncryptInterface(ctrl)
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
ctrl := gomock.NewController(t)

state := FlowStateCookie{}
sj, _ := json.Marshal(state)
mockLogger := NewMockLoggerInterface(ctrl)
mockEncrypt := NewMockEncryptInterface(ctrl)

mockEncrypt.EXPECT().Decrypt("mock-state").Return(string(sj), nil)
mockRequest := httptest.NewRequest(http.MethodGet, "/", nil)
mockRequest.AddCookie(&http.Cookie{Name: "state"})

mockRequest := httptest.NewRequest(http.MethodGet, "/", nil)
mockRequest.AddCookie(&http.Cookie{Name: "login_ui_state", Value: "mock-state"})
mockResponse := httptest.NewRecorder()

manager := NewAuthCookieManager(5, mockEncrypt, mockLogger)
cookie, err := manager.GetStateCookie(mockRequest)
manager := NewAuthCookieManager(5, mockEncrypt, mockLogger)
manager.ClearStateCookie(mockResponse)

if cookie != state {
t.Fatal("state cookie value does not match expected")
}

if err != nil {
t.Fatalf("expected error to be nil not %v", err)
}
}

func TestAuthCookieManager_GetStateCookieNoCookie(t *testing.T) {
ctrl := gomock.NewController(t)

mockLogger := NewMockLoggerInterface(ctrl)
mockRequest := httptest.NewRequest(http.MethodGet, "/", nil)

manager := NewAuthCookieManager(5, nil, mockLogger)
cookie, err := manager.GetStateCookie(mockRequest)

state := FlowStateCookie{}
if cookie != state {
t.Fatal("state cookie value does not match expected")
}
c, _ := findCookie("login_ui_state", mockResponse.Result().Cookies())

if err != nil {
t.Fatalf("expected error to be nil, not %v", err)
if c.Expires != epoch {
t.Fatal("did not clear state cookie")
}
})
}
}

func TestAuthCookieManager_GetStateCookieDecryptFailure(t *testing.T) {
ctrl := gomock.NewController(t)
mockError := errors.New("mock-error")

mockLogger := NewMockLoggerInterface(ctrl)
mockLogger.EXPECT().Errorf("can't decrypt cookie value, %v", mockError).Times(1)

mockEncrypt := NewMockEncryptInterface(ctrl)
mockEncrypt.EXPECT().Decrypt("mock-state").Return("", mockError)

mockRequest := httptest.NewRequest(http.MethodGet, "/", nil)
mockRequest.AddCookie(&http.Cookie{Name: "login_ui_state", Value: "mock-state"})

manager := NewAuthCookieManager(5, mockEncrypt, mockLogger)
cookie, err := manager.GetStateCookie(mockRequest)

state := FlowStateCookie{}
if cookie != state {
t.Fatal("state cookie value does not match expected")
func TestAuthCookieManager_GetStateCookie(t *testing.T) {
tests := []struct {
name string
setupMocks func(*MockEncryptInterface, *MockLoggerInterface)
requestCookie *http.Cookie
expectedCookie FlowStateCookie
expectedErr bool
}{
{
name: "Success",
setupMocks: func(mockEncrypt *MockEncryptInterface, mockLogger *MockLoggerInterface) {
state := FlowStateCookie{}
sj, _ := json.Marshal(state)
mockEncrypt.EXPECT().Decrypt("mock-state").Return(string(sj), nil)
},
requestCookie: &http.Cookie{Name: "login_ui_state", Value: "mock-state"},
expectedCookie: FlowStateCookie{},
expectedErr: false,
},
{
name: "NoCookie",
setupMocks: func(mockEncrypt *MockEncryptInterface, mockLogger *MockLoggerInterface) {},
requestCookie: nil,
expectedCookie: FlowStateCookie{},
expectedErr: false,
},
{
name: "DecryptFailure",
setupMocks: func(mockEncrypt *MockEncryptInterface, mockLogger *MockLoggerInterface) {
mockError := errors.New("mock-error")
mockLogger.EXPECT().Errorf("can't decrypt cookie value, %v", mockError).Times(1)
mockEncrypt.EXPECT().Decrypt("mock-state").Return("", mockError)
},
requestCookie: &http.Cookie{Name: "login_ui_state", Value: "mock-state"},
expectedCookie: FlowStateCookie{},
expectedErr: true,
},
}

if err == nil {
t.Fatalf("expected error to be not nil")
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
ctrl := gomock.NewController(t)

mockLogger := NewMockLoggerInterface(ctrl)
mockEncrypt := NewMockEncryptInterface(ctrl)

if tt.requestCookie == nil {
mockEncrypt = nil
}
if mockEncrypt != nil {
tt.setupMocks(mockEncrypt, mockLogger)
} else {
tt.setupMocks(nil, mockLogger)
}
Comment on lines +111 to +118

Copilot AI Feb 2, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The conditional logic reassigning mockEncrypt to nil and then passing it to setupMocks is convoluted. Consider having setupMocks handle nil cases directly, or restructure to avoid reassigning the mock variable after initialization.

Copilot uses AI. Check for mistakes.

mockRequest := httptest.NewRequest(http.MethodGet, "/", nil)
if tt.requestCookie != nil {
mockRequest.AddCookie(tt.requestCookie)
}

manager := NewAuthCookieManager(5, mockEncrypt, mockLogger)
cookie, err := manager.GetStateCookie(mockRequest)

if cookie != tt.expectedCookie {
t.Fatal("state cookie value does not match expected")
}

if tt.expectedErr {
if err == nil {
t.Fatalf("expected error to be not nil")
}
} else if err != nil {
t.Fatalf("expected error to be nil not %v", err)
}
})
}
}

func TestAuthCookieManager_SetStateCookie(t *testing.T) {
ctrl := gomock.NewController(t)

mockLogger := NewMockLoggerInterface(ctrl)
mockEncrypt := NewMockEncryptInterface(ctrl)

state := FlowStateCookie{}
js, _ := json.Marshal(state)

mockEncrypt.EXPECT().Encrypt(string(js)).Return("mock-state", nil)

mockResponse := httptest.NewRecorder()

manager := NewAuthCookieManager(5, mockEncrypt, mockLogger)
err := manager.SetStateCookie(mockResponse, state)

c, found := findCookie("login_ui_state", mockResponse.Result().Cookies())

if !found {
t.Fatal("did not set state cookie")
tests := []struct {
name string
setupMocks func(*MockEncryptInterface, *MockLoggerInterface)
expectedErr bool
}{
{
name: "Success",
setupMocks: func(mockEncrypt *MockEncryptInterface, mockLogger *MockLoggerInterface) {
state := FlowStateCookie{}
js, _ := json.Marshal(state)
mockEncrypt.EXPECT().Encrypt(string(js)).Return("mock-state", nil)
},
expectedErr: false,
},
{
name: "Failure",
setupMocks: func(mockEncrypt *MockEncryptInterface, mockLogger *MockLoggerInterface) {
mockError := errors.New("mock-error")
state := FlowStateCookie{}
js, _ := json.Marshal(state)
mockLogger.EXPECT().Errorf("can't encrypt cookie value, %v", mockError).Times(1)
mockEncrypt.EXPECT().Encrypt(string(js)).Return("", mockError)
},
expectedErr: true,
},
}

if c.Value != "mock-state" {
t.Fatal("state cookie value does not match expected")
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
ctrl := gomock.NewController(t)

if err != nil {
t.Fatalf("expected error to be nil not %v", err)
}
}
mockLogger := NewMockLoggerInterface(ctrl)
mockEncrypt := NewMockEncryptInterface(ctrl)

tt.setupMocks(mockEncrypt, mockLogger)

func TestAuthCookieManager_SetStateCookieFailure(t *testing.T) {
ctrl := gomock.NewController(t)
state := FlowStateCookie{}
mockResponse := httptest.NewRecorder()

mockError := errors.New("mock-error")
state := FlowStateCookie{}
js, _ := json.Marshal(state)
manager := NewAuthCookieManager(5, mockEncrypt, mockLogger)
err := manager.SetStateCookie(mockResponse, state)

mockLogger := NewMockLoggerInterface(ctrl)
mockLogger.EXPECT().Errorf("can't encrypt cookie value, %v", mockError).Times(1)
if tt.expectedErr {
if err == nil {
t.Fatalf("expected error to be not nil")
}
return
}

mockEncrypt := NewMockEncryptInterface(ctrl)
mockEncrypt.EXPECT().Encrypt(string(js)).Return("", mockError)
c, found := findCookie("login_ui_state", mockResponse.Result().Cookies())

mockResponse := httptest.NewRecorder()
if !found {
t.Fatal("did not set state cookie")
}

manager := NewAuthCookieManager(5, mockEncrypt, mockLogger)
err := manager.SetStateCookie(mockResponse, state)
if c.Value != "mock-state" {
t.Fatal("state cookie value does not match expected")
}

if err == nil {
t.Fatalf("expected error to be not nil")
if err != nil {
t.Fatalf("expected error to be nil not %v", err)
}
})
}
}
4 changes: 4 additions & 0 deletions pkg/kratos/encryption.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
// Copyright 2024 Canonical Ltd.
// SPDX-License-Identifier: AGPL-3.0

// Package kratos provides AES encryption and decryption utilities for securing flow state.
package kratos

import (
Expand Down Expand Up @@ -63,6 +64,7 @@ func (e *Encrypt) Decrypt(hexData string) (string, error) {
return string(decryptedData), nil
}

// splitNonceFromPayload extracts the nonce and payload from an encrypted byte array.
func (e *Encrypt) splitNonceFromPayload(encrypted []byte) ([]byte, []byte, error) {
nonceSize := e.gcm.NonceSize()
if len(encrypted) <= nonceSize {
Expand All @@ -73,6 +75,7 @@ func (e *Encrypt) splitNonceFromPayload(encrypted []byte) ([]byte, []byte, error
return noncePart, payloadPart, nil
}

// generateCipherNonce generates a random nonce for encryption.
func (e *Encrypt) generateCipherNonce() ([]byte, error) {
nonce := make([]byte, e.gcm.NonceSize())
if _, err := ioReadFull(rand.Reader, nonce); err != nil {
Expand All @@ -82,6 +85,7 @@ func (e *Encrypt) generateCipherNonce() ([]byte, error) {
return nonce, nil
}

// NewEncrypt creates a new Encrypt instance with the provided secret key and interfaces.
func NewEncrypt(secretKey []byte, logger logging.LoggerInterface, tracer tracing.TracingInterface) *Encrypt {
e := new(Encrypt)
c, err := aes.NewCipher(secretKey)
Expand Down
Loading
Loading