179 lines
5.6 KiB
Go
179 lines
5.6 KiB
Go
package httpx
|
|
|
|
import (
|
|
"encoding/json"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/stretchr/testify/require"
|
|
"senlinai-agent/backend/internal/config"
|
|
)
|
|
|
|
func TestHealthz(t *testing.T) {
|
|
router := NewRouter(config.Config{Env: "test"})
|
|
req := httptest.NewRequest(http.MethodGet, "/healthz", nil)
|
|
rec := httptest.NewRecorder()
|
|
|
|
router.ServeHTTP(rec, req)
|
|
|
|
require.Equal(t, http.StatusOK, rec.Code)
|
|
require.JSONEq(t, `{"status":"ok"}`, rec.Body.String())
|
|
}
|
|
|
|
func TestNewRouterRegistersFeatureRoutesUnderAPIV1(t *testing.T) {
|
|
router := NewRouter(config.Config{Env: "test"}, testRegistrar{})
|
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/ping", nil)
|
|
rec := httptest.NewRecorder()
|
|
|
|
router.ServeHTTP(rec, req)
|
|
|
|
require.Equal(t, http.StatusOK, rec.Code)
|
|
require.JSONEq(t, `{"pong":true}`, rec.Body.String())
|
|
}
|
|
|
|
func TestStatusLivesUnderAPIV1(t *testing.T) {
|
|
router := NewRouter(config.Config{Env: "test"})
|
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/status", nil)
|
|
recorder := httptest.NewRecorder()
|
|
|
|
router.ServeHTTP(recorder, req)
|
|
|
|
require.Equal(t, http.StatusOK, recorder.Code)
|
|
}
|
|
|
|
func TestLegacyAPIRouteIsNotExposed(t *testing.T) {
|
|
router := NewRouter(config.Config{Env: "test"}, testRegistrar{})
|
|
req := httptest.NewRequest(http.MethodGet, "/api/ping", nil)
|
|
recorder := httptest.NewRecorder()
|
|
|
|
router.ServeHTTP(recorder, req)
|
|
|
|
require.Equal(t, http.StatusNotFound, recorder.Code)
|
|
}
|
|
|
|
func TestAPIStatusReturnsChangingRFC3339UTCTimestampWithoutAuth(t *testing.T) {
|
|
router := NewProtectedRouter(config.Config{Env: "test"}, func(token string) (uint, error) {
|
|
return 0, http.ErrNoCookie
|
|
})
|
|
|
|
first := getStatusTimestamp(t, router)
|
|
time.Sleep(2 * time.Millisecond)
|
|
second := getStatusTimestamp(t, router)
|
|
|
|
require.Equal(t, time.UTC, first.Location())
|
|
require.Equal(t, time.UTC, second.Location())
|
|
require.True(t, second.After(first))
|
|
}
|
|
|
|
func TestRouterAddsCORSHeadersForLocalWebClient(t *testing.T) {
|
|
router := NewRouter(config.Config{Env: "test"}, testRegistrar{})
|
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/ping", nil)
|
|
req.Header.Set("Origin", "http://localhost:5173")
|
|
rec := httptest.NewRecorder()
|
|
|
|
router.ServeHTTP(rec, req)
|
|
|
|
require.Equal(t, http.StatusOK, rec.Code)
|
|
require.Equal(t, "http://localhost:5173", rec.Header().Get("Access-Control-Allow-Origin"))
|
|
require.Contains(t, rec.Header().Get("Access-Control-Allow-Headers"), "Authorization")
|
|
}
|
|
|
|
func TestRouterAddsCORSHeadersForMiniClient(t *testing.T) {
|
|
router := NewRouter(config.Config{Env: "test"}, testRegistrar{})
|
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/ping", nil)
|
|
req.Header.Set("Origin", "http://localhost:5180")
|
|
rec := httptest.NewRecorder()
|
|
|
|
router.ServeHTTP(rec, req)
|
|
|
|
require.Equal(t, http.StatusOK, rec.Code)
|
|
require.Equal(t, "http://localhost:5180", rec.Header().Get("Access-Control-Allow-Origin"))
|
|
}
|
|
|
|
func TestAPIV1LoginRemainsPublic(t *testing.T) {
|
|
router := NewProtectedRouter(config.Config{Env: "test"}, func(token string) (uint, error) {
|
|
return 0, http.ErrNoCookie
|
|
}, loginRegistrar{})
|
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/auth/login", nil)
|
|
rec := httptest.NewRecorder()
|
|
|
|
router.ServeHTTP(rec, req)
|
|
|
|
require.Equal(t, http.StatusOK, rec.Code)
|
|
}
|
|
|
|
func TestRouterUsesConfiguredCORSOrigins(t *testing.T) {
|
|
router := NewRouter(config.Config{Env: "test", AllowedOrigins: []string{"https://workbench.example.com"}}, testRegistrar{})
|
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/ping", nil)
|
|
req.Header.Set("Origin", "https://workbench.example.com")
|
|
rec := httptest.NewRecorder()
|
|
|
|
router.ServeHTTP(rec, req)
|
|
|
|
require.Equal(t, "https://workbench.example.com", rec.Header().Get("Access-Control-Allow-Origin"))
|
|
}
|
|
|
|
func TestRouterRejectsUnconfiguredCORSOrigin(t *testing.T) {
|
|
router := NewRouter(config.Config{Env: "production", AllowedOrigins: []string{"https://workbench.example.com"}}, testRegistrar{})
|
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/ping", nil)
|
|
req.Header.Set("Origin", "https://attacker.example.com")
|
|
rec := httptest.NewRecorder()
|
|
|
|
router.ServeHTTP(rec, req)
|
|
|
|
require.Empty(t, rec.Header().Get("Access-Control-Allow-Origin"))
|
|
}
|
|
|
|
func TestRouterHandlesCORSPreflightBeforeAuth(t *testing.T) {
|
|
router := NewProtectedRouter(config.Config{Env: "test"}, func(token string) (uint, error) {
|
|
return 0, http.ErrNoCookie
|
|
}, testRegistrar{})
|
|
req := httptest.NewRequest(http.MethodOptions, "/api/v1/ping", nil)
|
|
req.Header.Set("Origin", "http://localhost:5173")
|
|
req.Header.Set("Access-Control-Request-Method", "GET")
|
|
req.Header.Set("Access-Control-Request-Headers", "Authorization")
|
|
rec := httptest.NewRecorder()
|
|
|
|
router.ServeHTTP(rec, req)
|
|
|
|
require.Equal(t, http.StatusNoContent, rec.Code)
|
|
require.Equal(t, "http://localhost:5173", rec.Header().Get("Access-Control-Allow-Origin"))
|
|
require.Contains(t, rec.Header().Get("Access-Control-Allow-Headers"), "Authorization")
|
|
}
|
|
|
|
type testRegistrar struct{}
|
|
|
|
func (testRegistrar) Register(router gin.IRouter) {
|
|
router.GET("/ping", func(c *gin.Context) {
|
|
c.JSON(http.StatusOK, gin.H{"pong": true})
|
|
})
|
|
}
|
|
|
|
type loginRegistrar struct{}
|
|
|
|
func (loginRegistrar) Register(router gin.IRouter) {
|
|
router.POST("/auth/login", func(c *gin.Context) {
|
|
c.Status(http.StatusOK)
|
|
})
|
|
}
|
|
|
|
func getStatusTimestamp(t *testing.T, router http.Handler) time.Time {
|
|
t.Helper()
|
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/status", nil)
|
|
rec := httptest.NewRecorder()
|
|
|
|
router.ServeHTTP(rec, req)
|
|
|
|
require.Equal(t, http.StatusOK, rec.Code)
|
|
var payload struct {
|
|
Timestamp string `json:"timestamp"`
|
|
}
|
|
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload))
|
|
timestamp, err := time.Parse(time.RFC3339Nano, payload.Timestamp)
|
|
require.NoError(t, err)
|
|
return timestamp
|
|
}
|