File size: 3,730 Bytes
bf9e111
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
// Package middleware provides HTTP middleware components for the CLI Proxy API server.
// This file contains unit tests for the middleware components.
package middleware

import (
	"errors"
	"net/http"
	"net/http/httptest"
	"testing"

	"github.com/gin-gonic/gin"
	"github.com/stretchr/testify/assert"
)

func setupTestRouter() *gin.Engine {
	gin.SetMode(gin.TestMode)
	return gin.New()
}

func TestCorrelationIDMiddleware(t *testing.T) {
	router := setupTestRouter()
	router.Use(CorrelationIDMiddleware())
	router.GET("/test", func(c *gin.Context) {
		id := GetCorrelationID(c)
		c.JSON(200, gin.H{"correlation_id": id})
	})

	t.Run("generates correlation ID when not provided", func(t *testing.T) {
		w := httptest.NewRecorder()
		req, _ := http.NewRequest("GET", "/test", nil)
		router.ServeHTTP(w, req)

		assert.Equal(t, 200, w.Code)
		// Check that response contains a correlation ID
		assert.Contains(t, w.Body.String(), "correlation_id")
		// Check that response header contains the correlation ID
		assert.NotEmpty(t, w.Header().Get(CorrelationIDHeader))
	})

	t.Run("uses existing correlation ID from header", func(t *testing.T) {
		w := httptest.NewRecorder()
		req, _ := http.NewRequest("GET", "/test", nil)
		req.Header.Set(CorrelationIDHeader, "test-correlation-id-123")
		router.ServeHTTP(w, req)

		assert.Equal(t, 200, w.Code)
		assert.Contains(t, w.Body.String(), "test-correlation-id-123")
		assert.Equal(t, "test-correlation-id-123", w.Header().Get(CorrelationIDHeader))
	})
}

func TestErrorHandlerMiddleware(t *testing.T) {
	router := setupTestRouter()
	router.Use(CorrelationIDMiddleware())
	router.Use(ErrorHandlerMiddleware())

	router.GET("/error", func(c *gin.Context) {
		c.Error(errors.New("test error"))
		c.Abort()
	})

	router.GET("/success", func(c *gin.Context) {
		c.JSON(200, gin.H{"status": "ok"})
	})

	t.Run("handles errors gracefully", func(t *testing.T) {
		w := httptest.NewRecorder()
		req, _ := http.NewRequest("GET", "/error", nil)
		router.ServeHTTP(w, req)

		assert.Equal(t, 500, w.Code)
		assert.Contains(t, w.Body.String(), "INTERNAL_ERROR")
	})

	t.Run("passes through successful requests", func(t *testing.T) {
		w := httptest.NewRecorder()
		req, _ := http.NewRequest("GET", "/success", nil)
		router.ServeHTTP(w, req)

		assert.Equal(t, 200, w.Code)
		assert.Contains(t, w.Body.String(), "ok")
	})
}

func TestRecoveryMiddleware(t *testing.T) {
	router := setupTestRouter()
	router.Use(RecoveryMiddleware(nil))
	router.Use(CorrelationIDMiddleware())
	router.Use(ErrorHandlerMiddleware())

	router.GET("/panic", func(c *gin.Context) {
		panic("test panic")
	})

	router.GET("/normal", func(c *gin.Context) {
		c.JSON(200, gin.H{"status": "ok"})
	})

	t.Run("recovers from panic", func(t *testing.T) {
		w := httptest.NewRecorder()
		req, _ := http.NewRequest("GET", "/panic", nil)
		router.ServeHTTP(w, req)

		assert.Equal(t, 500, w.Code)
		assert.Contains(t, w.Body.String(), "INTERNAL_ERROR")
	})

	t.Run("normal requests work after panic recovery", func(t *testing.T) {
		w := httptest.NewRecorder()
		req, _ := http.NewRequest("GET", "/normal", nil)
		router.ServeHTTP(w, req)

		assert.Equal(t, 200, w.Code)
		assert.Contains(t, w.Body.String(), "ok")
	})
}

func TestSafeHandler(t *testing.T) {
	router := setupTestRouter()
	router.Use(CorrelationIDMiddleware())
	router.Use(ErrorHandlerMiddleware())

	router.GET("/safe-panic", SafeHandler(func(c *gin.Context) {
		panic("safe handler panic")
	}))

	t.Run("safe handler recovers from panic", func(t *testing.T) {
		w := httptest.NewRecorder()
		req, _ := http.NewRequest("GET", "/safe-panic", nil)
		router.ServeHTTP(w, req)

		assert.Equal(t, 500, w.Code)
		assert.Contains(t, w.Body.String(), "INTERNAL_ERROR")
	})
}