316 lines
8.7 KiB
Go
316 lines
8.7 KiB
Go
package proxy
|
|||
|
|
|
||
|
|
import (
|
||
|
|
"bytes"
|
||
|
|
"encoding/json"
|
||
|
|
"io"
|
||
|
|
"net/http"
|
||
|
|
"net/http/httptest"
|
||
|
|
"net/url"
|
||
|
|
"strings"
|
||
|
|
"testing"
|
||
|
|
|
||
|
|
"forgejo.riotpiao.com/rock/homelab-frontend/internal/config"
|
||
|
|
)
|
||
|
|
|
||
|
|
// stubAuthClient captures the form data sent and returns a canned response.
|
||
|
|
type stubAuthClient struct {
|
||
|
|
lastForm url.Values
|
||
|
|
statusCode int
|
||
|
|
body string
|
||
|
|
err error
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *stubAuthClient) PostForm(u string, data url.Values) (*http.Response, error) {
|
||
|
|
s.lastForm = data
|
||
|
|
if s.err != nil {
|
||
|
|
return nil, s.err
|
||
|
|
}
|
||
|
|
return &http.Response{
|
||
|
|
StatusCode: s.statusCode,
|
||
|
|
Body: io.NopCloser(strings.NewReader(s.body)),
|
||
|
|
Header: http.Header{"Content-Type": {"application/json"}},
|
||
|
|
}, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func newTestHandler(tokenURL, clientID, clientSecret string, client authClient) *Handler {
|
||
|
|
cfg := &config.Config{
|
||
|
|
Auth: config.AuthConfig{
|
||
|
|
TokenURL: tokenURL,
|
||
|
|
ClientID: clientID,
|
||
|
|
ClientSecret: clientSecret,
|
||
|
|
},
|
||
|
|
}
|
||
|
|
h := &Handler{
|
||
|
|
config: cfg,
|
||
|
|
authHTTP: client,
|
||
|
|
routes: make(map[string]*Route),
|
||
|
|
transports: make(map[string]*http.Transport),
|
||
|
|
}
|
||
|
|
return h
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAuthToken_Success(t *testing.T) {
|
||
|
|
stub := &stubAuthClient{
|
||
|
|
statusCode: 200,
|
||
|
|
body: `{"access_token":"jwt.token.here","refresh_token":"refresh","expires_in":3600}`,
|
||
|
|
}
|
||
|
|
h := newTestHandler("https://auth.example.com/token/", "api-gw", "secret", stub)
|
||
|
|
|
||
|
|
body := `{"username":"rock","password":"pass123"}`
|
||
|
|
r := httptest.NewRequest("POST", "/auth/token", strings.NewReader(body))
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthToken(w, r)
|
||
|
|
|
||
|
|
if w.Code != 200 {
|
||
|
|
t.Fatalf("expected 200, got %d: %s", w.Code, w.Body.String())
|
||
|
|
}
|
||
|
|
|
||
|
|
// Verify form sent to upstream
|
||
|
|
if stub.lastForm.Get("grant_type") != "password" {
|
||
|
|
t.Errorf("grant_type = %q, want password", stub.lastForm.Get("grant_type"))
|
||
|
|
}
|
||
|
|
if stub.lastForm.Get("username") != "rock" {
|
||
|
|
t.Errorf("username = %q, want rock", stub.lastForm.Get("username"))
|
||
|
|
}
|
||
|
|
if stub.lastForm.Get("client_id") != "api-gw" {
|
||
|
|
t.Errorf("client_id = %q, want api-gw", stub.lastForm.Get("client_id"))
|
||
|
|
}
|
||
|
|
if stub.lastForm.Get("client_secret") != "secret" {
|
||
|
|
t.Errorf("client_secret = %q, want secret", stub.lastForm.Get("client_secret"))
|
||
|
|
}
|
||
|
|
if stub.lastForm.Get("scope") != "openid roles permissions" {
|
||
|
|
t.Errorf("scope = %q, want default scope", stub.lastForm.Get("scope"))
|
||
|
|
}
|
||
|
|
|
||
|
|
// Verify response forwarded
|
||
|
|
var resp map[string]interface{}
|
||
|
|
json.NewDecoder(w.Body).Decode(&resp)
|
||
|
|
if resp["access_token"] != "jwt.token.here" {
|
||
|
|
t.Errorf("access_token not forwarded")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAuthToken_CustomScope(t *testing.T) {
|
||
|
|
stub := &stubAuthClient{statusCode: 200, body: `{}`}
|
||
|
|
h := newTestHandler("https://auth.example.com/token/", "api-gw", "secret", stub)
|
||
|
|
|
||
|
|
body := `{"username":"rock","password":"pass","scope":"openid roles"}`
|
||
|
|
r := httptest.NewRequest("POST", "/auth/token", strings.NewReader(body))
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthToken(w, r)
|
||
|
|
|
||
|
|
if stub.lastForm.Get("scope") != "openid roles" {
|
||
|
|
t.Errorf("scope = %q, want custom scope", stub.lastForm.Get("scope"))
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAuthToken_MissingUsername(t *testing.T) {
|
||
|
|
h := newTestHandler("https://auth.example.com/token/", "api-gw", "secret", &stubAuthClient{})
|
||
|
|
|
||
|
|
body := `{"password":"pass"}`
|
||
|
|
r := httptest.NewRequest("POST", "/auth/token", strings.NewReader(body))
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthToken(w, r)
|
||
|
|
|
||
|
|
if w.Code != 400 {
|
||
|
|
t.Errorf("expected 400, got %d", w.Code)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAuthToken_MissingPassword(t *testing.T) {
|
||
|
|
h := newTestHandler("https://auth.example.com/token/", "api-gw", "secret", &stubAuthClient{})
|
||
|
|
|
||
|
|
body := `{"username":"rock"}`
|
||
|
|
r := httptest.NewRequest("POST", "/auth/token", strings.NewReader(body))
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthToken(w, r)
|
||
|
|
|
||
|
|
if w.Code != 400 {
|
||
|
|
t.Errorf("expected 400, got %d", w.Code)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAuthToken_InvalidJSON(t *testing.T) {
|
||
|
|
h := newTestHandler("https://auth.example.com/token/", "api-gw", "secret", &stubAuthClient{})
|
||
|
|
|
||
|
|
r := httptest.NewRequest("POST", "/auth/token", strings.NewReader("not json"))
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthToken(w, r)
|
||
|
|
|
||
|
|
if w.Code != 400 {
|
||
|
|
t.Errorf("expected 400, got %d", w.Code)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAuthToken_WrongMethod(t *testing.T) {
|
||
|
|
h := newTestHandler("https://auth.example.com/token/", "api-gw", "secret", &stubAuthClient{})
|
||
|
|
|
||
|
|
r := httptest.NewRequest("GET", "/auth/token", nil)
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthToken(w, r)
|
||
|
|
|
||
|
|
if w.Code != 405 {
|
||
|
|
t.Errorf("expected 405, got %d", w.Code)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAuthToken_NotConfigured(t *testing.T) {
|
||
|
|
h := newTestHandler("", "", "", &stubAuthClient{})
|
||
|
|
|
||
|
|
body := `{"username":"rock","password":"pass"}`
|
||
|
|
r := httptest.NewRequest("POST", "/auth/token", strings.NewReader(body))
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthToken(w, r)
|
||
|
|
|
||
|
|
if w.Code != 503 {
|
||
|
|
t.Errorf("expected 503, got %d", w.Code)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAuthToken_UpstreamError(t *testing.T) {
|
||
|
|
stub := &stubAuthClient{err: io.ErrUnexpectedEOF}
|
||
|
|
h := newTestHandler("https://auth.example.com/token/", "api-gw", "secret", stub)
|
||
|
|
|
||
|
|
body := `{"username":"rock","password":"pass"}`
|
||
|
|
r := httptest.NewRequest("POST", "/auth/token", strings.NewReader(body))
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthToken(w, r)
|
||
|
|
|
||
|
|
if w.Code != 502 {
|
||
|
|
t.Errorf("expected 502, got %d", w.Code)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAuthToken_UpstreamRejectsCredentials(t *testing.T) {
|
||
|
|
stub := &stubAuthClient{
|
||
|
|
statusCode: 400,
|
||
|
|
body: `{"error":"invalid_grant","error_description":"bad password"}`,
|
||
|
|
}
|
||
|
|
h := newTestHandler("https://auth.example.com/token/", "api-gw", "secret", stub)
|
||
|
|
|
||
|
|
body := `{"username":"rock","password":"wrong"}`
|
||
|
|
r := httptest.NewRequest("POST", "/auth/token", strings.NewReader(body))
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthToken(w, r)
|
||
|
|
|
||
|
|
// Upstream error forwarded verbatim
|
||
|
|
if w.Code != 400 {
|
||
|
|
t.Errorf("expected 400 (forwarded), got %d", w.Code)
|
||
|
|
}
|
||
|
|
if !strings.Contains(w.Body.String(), "invalid_grant") {
|
||
|
|
t.Errorf("expected upstream error forwarded, got %s", w.Body.String())
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// --- /auth/refresh tests ---
|
||
|
|
|
||
|
|
func TestAuthRefresh_Success(t *testing.T) {
|
||
|
|
stub := &stubAuthClient{
|
||
|
|
statusCode: 200,
|
||
|
|
body: `{"access_token":"new.jwt","refresh_token":"new.refresh","expires_in":3600}`,
|
||
|
|
}
|
||
|
|
h := newTestHandler("https://auth.example.com/token/", "api-gw", "secret", stub)
|
||
|
|
|
||
|
|
body := `{"refresh_token":"old.refresh"}`
|
||
|
|
r := httptest.NewRequest("POST", "/auth/refresh", strings.NewReader(body))
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthRefresh(w, r)
|
||
|
|
|
||
|
|
if w.Code != 200 {
|
||
|
|
t.Fatalf("expected 200, got %d: %s", w.Code, w.Body.String())
|
||
|
|
}
|
||
|
|
if stub.lastForm.Get("grant_type") != "refresh_token" {
|
||
|
|
t.Errorf("grant_type = %q, want refresh_token", stub.lastForm.Get("grant_type"))
|
||
|
|
}
|
||
|
|
if stub.lastForm.Get("refresh_token") != "old.refresh" {
|
||
|
|
t.Errorf("refresh_token not sent")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAuthRefresh_MissingToken(t *testing.T) {
|
||
|
|
h := newTestHandler("https://auth.example.com/token/", "api-gw", "secret", &stubAuthClient{})
|
||
|
|
|
||
|
|
r := httptest.NewRequest("POST", "/auth/refresh", strings.NewReader(`{}`))
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthRefresh(w, r)
|
||
|
|
|
||
|
|
if w.Code != 400 {
|
||
|
|
t.Errorf("expected 400, got %d", w.Code)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAuthRefresh_WrongMethod(t *testing.T) {
|
||
|
|
h := newTestHandler("https://auth.example.com/token/", "api-gw", "secret", &stubAuthClient{})
|
||
|
|
|
||
|
|
r := httptest.NewRequest("GET", "/auth/refresh", nil)
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthRefresh(w, r)
|
||
|
|
|
||
|
|
if w.Code != 405 {
|
||
|
|
t.Errorf("expected 405, got %d", w.Code)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAuthRefresh_NotConfigured(t *testing.T) {
|
||
|
|
h := newTestHandler("", "", "", &stubAuthClient{})
|
||
|
|
|
||
|
|
r := httptest.NewRequest("POST", "/auth/refresh", strings.NewReader(`{"refresh_token":"x"}`))
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthRefresh(w, r)
|
||
|
|
|
||
|
|
if w.Code != 503 {
|
||
|
|
t.Errorf("expected 503, got %d", w.Code)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAuthRefresh_ExpiredToken(t *testing.T) {
|
||
|
|
stub := &stubAuthClient{
|
||
|
|
statusCode: 401,
|
||
|
|
body: `{"error":"invalid_grant","error_description":"token expired"}`,
|
||
|
|
}
|
||
|
|
h := newTestHandler("https://auth.example.com/token/", "api-gw", "secret", stub)
|
||
|
|
|
||
|
|
r := httptest.NewRequest("POST", "/auth/refresh", strings.NewReader(`{"refresh_token":"expired"}`))
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthRefresh(w, r)
|
||
|
|
|
||
|
|
if w.Code != 401 {
|
||
|
|
t.Errorf("expected 401 (forwarded), got %d", w.Code)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// Verify no credentials are leaked in response bodies
|
||
|
|
func TestAuthToken_NoCredentialLeak(t *testing.T) {
|
||
|
|
stub := &stubAuthClient{statusCode: 200, body: `{"access_token":"tok"}`}
|
||
|
|
h := newTestHandler("https://auth.example.com/token/", "api-gw", "super-secret", stub)
|
||
|
|
|
||
|
|
body := `{"username":"rock","password":"my-password"}`
|
||
|
|
r := httptest.NewRequest("POST", "/auth/token", bytes.NewReader([]byte(body)))
|
||
|
|
w := httptest.NewRecorder()
|
||
|
|
|
||
|
|
h.handleAuthToken(w, r)
|
||
|
|
|
||
|
|
respBody := w.Body.String()
|
||
|
|
if strings.Contains(respBody, "super-secret") {
|
||
|
|
t.Error("client_secret leaked in response")
|
||
|
|
}
|
||
|
|
if strings.Contains(respBody, "my-password") {
|
||
|
|
t.Error("password leaked in response")
|
||
|
|
}
|
||
|
|
}
|