83 lines
2.1 KiB
Go
83 lines
2.1 KiB
Go
package requestid
|
|
|
|
import (
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
|
|
"github.com/gin-gonic/gin"
|
|
)
|
|
|
|
func TestMiddlewareClientProvidedRequestID(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
var (
|
|
gotID string
|
|
gotProvided bool
|
|
)
|
|
router := gin.New()
|
|
router.Use(New(WithGenerator(func() string { return "generated-id" })))
|
|
router.GET("/", func(c *gin.Context) {
|
|
gotID, gotProvided = FromClient(c)
|
|
c.Status(http.StatusNoContent)
|
|
})
|
|
|
|
request := httptest.NewRequest(http.MethodGet, "/", nil)
|
|
request.Header.Set(HeaderKey, " client-id ")
|
|
response := httptest.NewRecorder()
|
|
router.ServeHTTP(response, request)
|
|
|
|
if gotID != "client-id" || !gotProvided {
|
|
t.Fatalf("FromClient() = (%q, %v), want (client-id, true)", gotID, gotProvided)
|
|
}
|
|
if got := response.Header().Get(HeaderKey); got != "client-id" {
|
|
t.Fatalf("response %s = %q, want client-id", HeaderKey, got)
|
|
}
|
|
}
|
|
|
|
func TestMiddlewareGeneratedRequestIDIsTracingOnly(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
tests := []struct {
|
|
name string
|
|
requestID string
|
|
}{
|
|
{name: "missing"},
|
|
{name: "blank", requestID: " "},
|
|
{name: "too long", requestID: strings.Repeat("x", 129)},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
var (
|
|
gotID string
|
|
gotProvided bool
|
|
)
|
|
router := gin.New()
|
|
router.Use(New(WithGenerator(func() string { return "generated-id" })))
|
|
router.GET("/", func(c *gin.Context) {
|
|
gotID, gotProvided = FromClient(c)
|
|
c.Status(http.StatusNoContent)
|
|
})
|
|
|
|
request := httptest.NewRequest(http.MethodGet, "/", nil)
|
|
if tt.requestID != "" {
|
|
request.Header.Set(HeaderKey, tt.requestID)
|
|
}
|
|
response := httptest.NewRecorder()
|
|
router.ServeHTTP(response, request)
|
|
|
|
if gotID != "" || gotProvided {
|
|
t.Fatalf("FromClient() = (%q, %v), want empty and false", gotID, gotProvided)
|
|
}
|
|
if got := request.Header.Get(HeaderKey); got != "generated-id" {
|
|
t.Fatalf("request %s = %q, want generated-id", HeaderKey, got)
|
|
}
|
|
if got := response.Header().Get(HeaderKey); got != "generated-id" {
|
|
t.Fatalf("response %s = %q, want generated-id", HeaderKey, got)
|
|
}
|
|
})
|
|
}
|
|
}
|