mirror of
https://github.com/h44z/wg-portal.git
synced 2026-10-08 06:26:41 +00:00
* feat: add link-based configuration emails (#725) * fix tests
This commit is contained in:
@@ -1,18 +0,0 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Base64UrlDecode decodes a base64 url encoded string.
|
||||
// In comparison to the standard base64 encoding, the url encoding uses - instead of + and _ instead of /
|
||||
// as well as . instead of =.
|
||||
func Base64UrlDecode(in string) string {
|
||||
in = strings.ReplaceAll(in, "-", "=")
|
||||
in = strings.ReplaceAll(in, "_", "/")
|
||||
in = strings.ReplaceAll(in, ".", "+")
|
||||
|
||||
output, _ := base64.StdEncoding.DecodeString(in)
|
||||
return string(output)
|
||||
}
|
||||
@@ -383,12 +383,6 @@ func (e AuthEndpoint) setAuthenticatedUser(r *http.Request, user *domain.User, o
|
||||
// @Router /auth/login [post]
|
||||
func (e AuthEndpoint) handleLoginPost() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
currentSession := e.session.GetData(r.Context())
|
||||
if currentSession.LoggedIn {
|
||||
respond.JSON(w, http.StatusOK, model.Error{Code: http.StatusOK, Message: "already logged in"})
|
||||
return
|
||||
}
|
||||
|
||||
var loginData struct {
|
||||
Username string `json:"username" binding:"required,min=2"`
|
||||
Password string `json:"password" binding:"required,min=4"`
|
||||
@@ -570,7 +564,7 @@ func (e AuthEndpoint) handleWebAuthnCredentialsDelete() http.HandlerFunc {
|
||||
|
||||
userIdentifier := domain.UserIdentifier(currentSession.UserIdentifier)
|
||||
|
||||
credentialId := Base64UrlDecode(request.Path(r, "id"))
|
||||
credentialId := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
|
||||
credentials, err := e.webAuthn.RemoveCredential(r.Context(), userIdentifier, credentialId)
|
||||
if err != nil {
|
||||
@@ -605,7 +599,7 @@ func (e AuthEndpoint) handleWebAuthnCredentialsPut() http.HandlerFunc {
|
||||
|
||||
userIdentifier := domain.UserIdentifier(currentSession.UserIdentifier)
|
||||
|
||||
credentialId := Base64UrlDecode(request.Path(r, "id"))
|
||||
credentialId := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
var req model.WebAuthnCredentialRequest
|
||||
if err := request.BodyJson(r, &req); err != nil {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
|
||||
@@ -2,11 +2,14 @@ package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/h44z/wg-portal/internal/config"
|
||||
"github.com/h44z/wg-portal/internal/domain"
|
||||
)
|
||||
|
||||
type testSession struct {
|
||||
@@ -115,3 +118,55 @@ func TestAuthEndpointFrontendUrlUsesBasePathAppMount(t *testing.T) {
|
||||
t.Fatalf("expected frontend URL %q, got %q", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
type dummyAuthService struct {
|
||||
loginErr error
|
||||
user *domain.User
|
||||
}
|
||||
|
||||
func (d dummyAuthService) GetExternalLoginProviders(_ context.Context) []domain.LoginProviderInfo {
|
||||
return nil
|
||||
}
|
||||
func (d dummyAuthService) PlainLogin(_ context.Context, username, password string) (*domain.User, error) {
|
||||
if d.loginErr != nil {
|
||||
return nil, d.loginErr
|
||||
}
|
||||
return d.user, nil
|
||||
}
|
||||
func (d dummyAuthService) OauthLoginStep1(_ context.Context, _ string) (string, string, string, string, error) {
|
||||
return "", "", "", "", nil
|
||||
}
|
||||
func (d dummyAuthService) OauthLoginStep2(_ context.Context, _, _, _, _ string) (*domain.User, string, error) {
|
||||
return nil, "", nil
|
||||
}
|
||||
func (d dummyAuthService) OauthProviderLogoutUrl(_, _, _ string) (string, bool) {
|
||||
return "", false
|
||||
}
|
||||
|
||||
type dummyValidator struct{}
|
||||
|
||||
func (d dummyValidator) Struct(_ any) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestAuthEndpointHandleLoginPostRejectsInvalidCredentialsEvenIfSessionDirty(t *testing.T) {
|
||||
session := &testSession{data: SessionData{
|
||||
LoggedIn: true,
|
||||
UserIdentifier: "previous-user",
|
||||
}}
|
||||
ep := AuthEndpoint{
|
||||
session: session,
|
||||
authService: dummyAuthService{loginErr: errors.New("auth failed")},
|
||||
validate: dummyValidator{},
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v0/auth/login", strings.NewReader(`{"username":"admin","password":"wrongpassword"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
ep.handleLoginPost().ServeHTTP(res, req)
|
||||
|
||||
if res.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("expected status %d (Unauthorized), got %d", http.StatusUnauthorized, res.Code)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -138,7 +138,7 @@ func (e InterfaceEndpoint) handleAllGet() http.HandlerFunc {
|
||||
// @Router /interface/get/{id} [get]
|
||||
func (e InterfaceEndpoint) handleSingleGet() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest, model.Error{
|
||||
Code: http.StatusInternalServerError, Message: "missing id parameter",
|
||||
@@ -170,7 +170,7 @@ func (e InterfaceEndpoint) handleSingleGet() http.HandlerFunc {
|
||||
// @Router /interface/config/{id} [get]
|
||||
func (e InterfaceEndpoint) handleConfigGet() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest, model.Error{
|
||||
Code: http.StatusInternalServerError, Message: "missing id parameter",
|
||||
@@ -212,7 +212,7 @@ func (e InterfaceEndpoint) handleConfigGet() http.HandlerFunc {
|
||||
// @Router /interface/{id} [put]
|
||||
func (e InterfaceEndpoint) handleUpdatePut() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusBadRequest, Message: "missing interface id"})
|
||||
@@ -293,7 +293,7 @@ func (e InterfaceEndpoint) handleCreatePost() http.HandlerFunc {
|
||||
// @Router /interface/peers/{id} [get]
|
||||
func (e InterfaceEndpoint) handlePeersGet() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest, model.Error{
|
||||
Code: http.StatusInternalServerError, Message: "missing id parameter",
|
||||
@@ -326,7 +326,7 @@ func (e InterfaceEndpoint) handlePeersGet() http.HandlerFunc {
|
||||
// @Router /interface/{id} [delete]
|
||||
func (e InterfaceEndpoint) handleDelete() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusBadRequest, Message: "missing interface id"})
|
||||
@@ -358,7 +358,7 @@ func (e InterfaceEndpoint) handleDelete() http.HandlerFunc {
|
||||
// @Router /interface/{id}/save-config [post]
|
||||
func (e InterfaceEndpoint) handleSaveConfigPost() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusBadRequest, Message: "missing interface id"})
|
||||
@@ -391,7 +391,7 @@ func (e InterfaceEndpoint) handleSaveConfigPost() http.HandlerFunc {
|
||||
// @Router /interface/{id}/apply-peer-defaults [post]
|
||||
func (e InterfaceEndpoint) handleApplyPeerDefaultsPost() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusBadRequest, Message: "missing interface id"})
|
||||
@@ -438,7 +438,7 @@ func (e InterfaceEndpoint) handleApplyPeerDefaultsPost() http.HandlerFunc {
|
||||
// @Router /interface/{id}/create-default-peers [post]
|
||||
func (e InterfaceEndpoint) handleCreateDefaultPeersPost() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusBadRequest, Message: "missing interface id"})
|
||||
|
||||
@@ -107,7 +107,7 @@ func (e PeerEndpoint) RegisterRoutes(g *routegroup.Bundle) {
|
||||
// @Router /peer/iface/{iface}/all [get]
|
||||
func (e PeerEndpoint) handleAllGet() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
interfaceId := Base64UrlDecode(request.Path(r, "iface"))
|
||||
interfaceId := domain.Base64UrlDecode(request.Path(r, "iface"))
|
||||
if interfaceId == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusBadRequest, Message: "missing iface parameter"})
|
||||
@@ -138,7 +138,7 @@ func (e PeerEndpoint) handleAllGet() http.HandlerFunc {
|
||||
// @Router /peer/{id} [get]
|
||||
func (e PeerEndpoint) handleSingleGet() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
peerId := Base64UrlDecode(request.Path(r, "id"))
|
||||
peerId := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if peerId == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusBadRequest, Message: "missing id parameter"})
|
||||
@@ -169,7 +169,7 @@ func (e PeerEndpoint) handleSingleGet() http.HandlerFunc {
|
||||
// @Router /peer/iface/{iface}/prepare [get]
|
||||
func (e PeerEndpoint) handlePrepareGet() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
interfaceId := Base64UrlDecode(request.Path(r, "iface"))
|
||||
interfaceId := domain.Base64UrlDecode(request.Path(r, "iface"))
|
||||
if interfaceId == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusBadRequest, Message: "missing iface parameter"})
|
||||
@@ -201,7 +201,7 @@ func (e PeerEndpoint) handlePrepareGet() http.HandlerFunc {
|
||||
// @Router /peer/iface/{iface}/new [post]
|
||||
func (e PeerEndpoint) handleCreatePost() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
interfaceId := Base64UrlDecode(request.Path(r, "iface"))
|
||||
interfaceId := domain.Base64UrlDecode(request.Path(r, "iface"))
|
||||
if interfaceId == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusBadRequest, Message: "missing iface parameter"})
|
||||
@@ -249,7 +249,7 @@ func (e PeerEndpoint) handleCreatePost() http.HandlerFunc {
|
||||
// @Router /peer/iface/{iface}/multiplenew [post]
|
||||
func (e PeerEndpoint) handleCreateMultiplePost() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
interfaceId := Base64UrlDecode(request.Path(r, "iface"))
|
||||
interfaceId := domain.Base64UrlDecode(request.Path(r, "iface"))
|
||||
if interfaceId == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusBadRequest, Message: "missing iface parameter"})
|
||||
@@ -292,7 +292,7 @@ func (e PeerEndpoint) handleCreateMultiplePost() http.HandlerFunc {
|
||||
// @Router /peer/{id} [put]
|
||||
func (e PeerEndpoint) handleUpdatePut() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
peerId := Base64UrlDecode(request.Path(r, "id"))
|
||||
peerId := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if peerId == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusBadRequest, Message: "missing id parameter"})
|
||||
@@ -339,7 +339,7 @@ func (e PeerEndpoint) handleUpdatePut() http.HandlerFunc {
|
||||
// @Router /peer/{id} [delete]
|
||||
func (e PeerEndpoint) handleDelete() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest, model.Error{Code: http.StatusBadRequest, Message: "missing peer id"})
|
||||
return
|
||||
@@ -370,7 +370,7 @@ func (e PeerEndpoint) handleDelete() http.HandlerFunc {
|
||||
// @Router /peer/config/{id} [get]
|
||||
func (e PeerEndpoint) handleConfigGet() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest, model.Error{
|
||||
Code: http.StatusInternalServerError, Message: "missing id parameter",
|
||||
@@ -415,7 +415,7 @@ func (e PeerEndpoint) handleConfigGet() http.HandlerFunc {
|
||||
// @Router /peer/config-qr/{id} [get]
|
||||
func (e PeerEndpoint) handleQrCodeGet() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest, model.Error{
|
||||
Code: http.StatusInternalServerError, Message: "missing id parameter",
|
||||
@@ -504,7 +504,7 @@ func (e PeerEndpoint) handleEmailPost() http.HandlerFunc {
|
||||
// @Router /peer/iface/{iface}/stats [get]
|
||||
func (e PeerEndpoint) handleStatsGet() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
interfaceId := Base64UrlDecode(request.Path(r, "iface"))
|
||||
interfaceId := domain.Base64UrlDecode(request.Path(r, "iface"))
|
||||
if interfaceId == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusBadRequest, Message: "missing iface parameter"})
|
||||
|
||||
@@ -125,7 +125,7 @@ func (e UserEndpoint) handleAllGet() http.HandlerFunc {
|
||||
// @Router /user/{id} [get]
|
||||
func (e UserEndpoint) handleSingleGet() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest, model.Error{Code: http.StatusBadRequest, Message: "missing user id"})
|
||||
return
|
||||
@@ -156,7 +156,7 @@ func (e UserEndpoint) handleSingleGet() http.HandlerFunc {
|
||||
// @Router /user/{id} [put]
|
||||
func (e UserEndpoint) handleUpdatePut() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest, model.Error{Code: http.StatusBadRequest, Message: "missing user id"})
|
||||
return
|
||||
@@ -236,7 +236,7 @@ func (e UserEndpoint) handleCreatePost() http.HandlerFunc {
|
||||
// @Router /user/{id}/peers [get]
|
||||
func (e UserEndpoint) handlePeersGet() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
userId := Base64UrlDecode(request.Path(r, "id"))
|
||||
userId := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if userId == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusInternalServerError, Message: "missing id parameter"})
|
||||
@@ -267,7 +267,7 @@ func (e UserEndpoint) handlePeersGet() http.HandlerFunc {
|
||||
// @Router /user/{id}/stats [get]
|
||||
func (e UserEndpoint) handleStatsGet() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
userId := Base64UrlDecode(request.Path(r, "id"))
|
||||
userId := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if userId == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusInternalServerError, Message: "missing id parameter"})
|
||||
@@ -298,7 +298,7 @@ func (e UserEndpoint) handleStatsGet() http.HandlerFunc {
|
||||
// @Router /user/{id}/interfaces [get]
|
||||
func (e UserEndpoint) handleInterfacesGet() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
userId := Base64UrlDecode(request.Path(r, "id"))
|
||||
userId := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if userId == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusInternalServerError, Message: "missing id parameter"})
|
||||
@@ -329,7 +329,7 @@ func (e UserEndpoint) handleInterfacesGet() http.HandlerFunc {
|
||||
// @Router /user/{id} [delete]
|
||||
func (e UserEndpoint) handleDelete() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id := Base64UrlDecode(request.Path(r, "id"))
|
||||
id := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if id == "" {
|
||||
respond.JSON(w, http.StatusBadRequest, model.Error{Code: http.StatusBadRequest, Message: "missing user id"})
|
||||
return
|
||||
@@ -358,7 +358,7 @@ func (e UserEndpoint) handleDelete() http.HandlerFunc {
|
||||
// @Router /user/{id}/api/enable [post]
|
||||
func (e UserEndpoint) handleApiEnablePost() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
userId := Base64UrlDecode(request.Path(r, "id"))
|
||||
userId := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if userId == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusInternalServerError, Message: "missing id parameter"})
|
||||
@@ -388,7 +388,7 @@ func (e UserEndpoint) handleApiEnablePost() http.HandlerFunc {
|
||||
// @Router /user/{id}/api/disable [post]
|
||||
func (e UserEndpoint) handleApiDisablePost() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
userId := Base64UrlDecode(request.Path(r, "id"))
|
||||
userId := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if userId == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusInternalServerError, Message: "missing id parameter"})
|
||||
@@ -418,7 +418,7 @@ func (e UserEndpoint) handleApiDisablePost() http.HandlerFunc {
|
||||
// @Router /user/{id}/change-password [post]
|
||||
func (e UserEndpoint) handleChangePasswordPost() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
userId := Base64UrlDecode(request.Path(r, "id"))
|
||||
userId := domain.Base64UrlDecode(request.Path(r, "id"))
|
||||
if userId == "" {
|
||||
respond.JSON(w, http.StatusBadRequest,
|
||||
model.Error{Code: http.StatusInternalServerError, Message: "missing id parameter"})
|
||||
|
||||
@@ -110,7 +110,7 @@ func (h AuthenticationHandler) UserIdMatch(idParameter string) func(next http.Ha
|
||||
}
|
||||
|
||||
sessionUserId := domain.UserIdentifier(session.UserIdentifier)
|
||||
requestUserId := domain.UserIdentifier(Base64UrlDecode(request.Path(r, idParameter)))
|
||||
requestUserId := domain.UserIdentifier(domain.Base64UrlDecode(request.Path(r, idParameter)))
|
||||
|
||||
if sessionUserId != requestUserId {
|
||||
// Abort the request with the appropriate error code
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/mail"
|
||||
"net/url"
|
||||
|
||||
"github.com/h44z/wg-portal/internal/config"
|
||||
"github.com/h44z/wg-portal/internal/domain"
|
||||
@@ -135,7 +136,8 @@ func (m Manager) sendPeerEmail(
|
||||
mailOptions domain.MailOptions
|
||||
)
|
||||
if linkOnly {
|
||||
txtMail, htmlMail, err = m.tplHandler.GetConfigMail(user, "deep link TBD")
|
||||
configDownloadLink := m.getPeerConfigDownloadLink(peer.Identifier, style)
|
||||
txtMail, htmlMail, err = m.tplHandler.GetConfigMail(user, configDownloadLink)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get mail body: %w", err)
|
||||
}
|
||||
@@ -182,6 +184,22 @@ func (m Manager) sendPeerEmail(
|
||||
return nil
|
||||
}
|
||||
|
||||
// getPeerConfigDownloadLink builds an absolute link that points to the peer configuration download
|
||||
// page of the WireGuard Portal web frontend. The link is used in link-only emails.
|
||||
//
|
||||
// The link is a "deep link" into the single-page application (hash based routing). When the recipient
|
||||
// opens the link while not being authenticated, the frontend redirects them to the login page first and
|
||||
// only starts the configuration download after a successful authentication.
|
||||
func (m Manager) getPeerConfigDownloadLink(peerId domain.PeerIdentifier, style string) string {
|
||||
encodedId := domain.Base64UrlEncode(string(peerId))
|
||||
link := fmt.Sprintf("%s%s/app/#/peer/config/%s",
|
||||
m.cfg.Web.ExternalUrl, m.cfg.Web.BasePath, encodedId)
|
||||
if style != "" {
|
||||
link += "?style=" + url.QueryEscape(style)
|
||||
}
|
||||
return link
|
||||
}
|
||||
|
||||
func (m Manager) resolveEmail(ctx context.Context, peer *domain.Peer) (string, domain.User) {
|
||||
user, err := m.users.GetUser(ctx, peer.UserIdentifier)
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
package mail
|
||||
|
||||
import (
|
||||
"io"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/h44z/wg-portal/internal/config"
|
||||
"github.com/h44z/wg-portal/internal/domain"
|
||||
)
|
||||
|
||||
func Test_base64UrlEncode_isReversibleWithHandlerDecode(t *testing.T) {
|
||||
inputs := []string{
|
||||
"peer-identifier",
|
||||
"aGVsbG8=", // ensure padding characters are handled
|
||||
"abc/def+ghi", // ensure + and / are handled
|
||||
"wgTestKey1234567890==",
|
||||
}
|
||||
|
||||
for _, in := range inputs {
|
||||
encoded := domain.Base64UrlEncode(in)
|
||||
|
||||
// The URL-safe variant must not contain characters that are unsafe in URLs.
|
||||
if strings.ContainsAny(encoded, "+/=") {
|
||||
t.Fatalf("encoded value %q still contains unsafe characters", encoded)
|
||||
}
|
||||
|
||||
decoded := domain.Base64UrlDecode(encoded)
|
||||
if decoded != in {
|
||||
t.Fatalf("round trip failed: got %q, want %q (encoded: %q)", decoded, in, encoded)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func Test_getPeerConfigDownloadLink(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
externalUrl string
|
||||
basePath string
|
||||
peerId domain.PeerIdentifier
|
||||
style string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "no base path",
|
||||
externalUrl: "https://wg.example.com",
|
||||
basePath: "",
|
||||
peerId: "peer1",
|
||||
style: "wgquick",
|
||||
want: "https://wg.example.com/app/#/peer/config/" + domain.Base64UrlEncode("peer1") + "?style=wgquick",
|
||||
},
|
||||
{
|
||||
name: "with base path",
|
||||
externalUrl: "https://wg.example.com",
|
||||
basePath: "/wg",
|
||||
peerId: "peer1",
|
||||
style: "",
|
||||
want: "https://wg.example.com/wg/app/#/peer/config/" + domain.Base64UrlEncode("peer1"),
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
cfg := &config.Config{}
|
||||
cfg.Web.ExternalUrl = tt.externalUrl
|
||||
cfg.Web.BasePath = tt.basePath
|
||||
m := Manager{cfg: cfg}
|
||||
|
||||
got := m.getPeerConfigDownloadLink(tt.peerId, tt.style)
|
||||
if got != tt.want {
|
||||
t.Fatalf("getPeerConfigDownloadLink() = %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func Test_GetConfigMail_containsLink(t *testing.T) {
|
||||
handler, err := newTemplateHandler("https://wg.example.com", "WireGuard Portal", "")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create template handler: %v", err)
|
||||
}
|
||||
|
||||
link := "https://wg.example.com/app/#/peer/config/abc?style=wgquick"
|
||||
txtReader, htmlReader, err := handler.GetConfigMail(&domain.User{Firstname: "John", Lastname: "Doe"}, link)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to render link mail: %v", err)
|
||||
}
|
||||
|
||||
txt, _ := io.ReadAll(txtReader)
|
||||
html, _ := io.ReadAll(htmlReader)
|
||||
|
||||
if !strings.Contains(string(txt), link) {
|
||||
t.Errorf("text link mail does not contain the generated link.\n%s", string(txt))
|
||||
}
|
||||
if !strings.Contains(string(html), link) {
|
||||
t.Errorf("html link mail does not contain the generated link.\n%s", string(html))
|
||||
}
|
||||
|
||||
// The link mail must not reference the placeholder that was used before the fix.
|
||||
if strings.Contains(string(txt), "deep link TBD") || strings.Contains(string(html), "deep link TBD") {
|
||||
t.Errorf("link mail still contains the placeholder link")
|
||||
}
|
||||
}
|
||||
@@ -80,7 +80,7 @@
|
||||
<tr>
|
||||
<td class="td container" style="width:650px; min-width:650px; font-size:0pt; line-height:0pt; margin:0; font-weight:normal; padding:55px 0px;">
|
||||
|
||||
<!-- Article / Image On The Left - Copy On The Right -->
|
||||
<!-- Article / Copy + Download Link -->
|
||||
<table width="100%" border="0" cellspacing="0" cellpadding="0">
|
||||
<tr>
|
||||
<td style="padding-bottom: 10px;">
|
||||
@@ -89,28 +89,28 @@
|
||||
<td class="tbrr p30-15" style="padding: 60px 30px; border-radius:26px 26px 0px 0px;" bgcolor="#ffffff">
|
||||
<table width="100%" border="0" cellspacing="0" cellpadding="0">
|
||||
<tr>
|
||||
<th class="column-top" width="210" style="font-size:0pt; line-height:0pt; padding:0; margin:0; font-weight:normal; vertical-align:top;">
|
||||
<table width="100%" border="0" cellspacing="0" cellpadding="0">
|
||||
{{if $.User.Firstname}}
|
||||
<td class="h4 pb20" style="color:#000000; font-family:'Muli', Arial,sans-serif; font-size:20px; line-height:28px; text-align:left; padding-bottom:20px;">Hello {{$.User.Firstname}} {{$.User.Lastname}}</td>
|
||||
{{else}}
|
||||
<td class="h4 pb20" style="color:#000000; font-family:'Muli', Arial,sans-serif; font-size:20px; line-height:28px; text-align:left; padding-bottom:20px;">Hello</td>
|
||||
{{end}}
|
||||
</tr>
|
||||
<tr>
|
||||
<td class="text pb20" style="color:#000000; font-family:Arial,sans-serif; font-size:14px; line-height:26px; text-align:left; padding-bottom:20px;">You or your administrator probably requested this VPN configuration. Use the button below to download your personal WireGuard configuration file and open it in the WireGuard VPN client to establish a secure VPN connection.</td>
|
||||
</tr>
|
||||
<!-- Button -->
|
||||
<tr>
|
||||
<td align="left">
|
||||
<table border="0" cellspacing="0" cellpadding="0">
|
||||
<tr>
|
||||
<td class="fluid-img" style="font-size:0pt; line-height:0pt; text-align:left;"><img src="cid:{{$.QrcodePngName}}" width="210" height="210" border="0" alt="" /></td>
|
||||
<td class="blue-button text-button" style="background:#000000; color:#ffffff; font-family:'Muli', Arial,sans-serif; font-size:14px; line-height:18px; padding:12px 30px; text-align:center; border-radius:0px 22px 22px 22px; font-weight:bold;"><a href="{{$.Link}}" target="_blank" rel="noopener noreferrer" class="link-white" style="color:#ffffff; text-decoration:none;"><span class="link-white" style="color:#ffffff; text-decoration:none;">Download VPN Configuration</span></a></td>
|
||||
</tr>
|
||||
</table>
|
||||
</th>
|
||||
<th class="column-empty2" width="30" style="font-size:0pt; line-height:0pt; padding:0; margin:0; font-weight:normal; vertical-align:top;"></th>
|
||||
<th class="column-top" width="280" style="font-size:0pt; line-height:0pt; padding:0; margin:0; font-weight:normal; vertical-align:top;">
|
||||
<table width="100%" border="0" cellspacing="0" cellpadding="0">
|
||||
<tr>
|
||||
{{if $.User.Firstname}}
|
||||
<td class="h4 pb20" style="color:#000000; font-family:'Muli', Arial,sans-serif; font-size:20px; line-height:28px; text-align:left; padding-bottom:20px;">Hello {{$.User.Firstname}} {{$.User.Lastname}}</td>
|
||||
{{else}}
|
||||
<td class="h4 pb20" style="color:#000000; font-family:'Muli', Arial,sans-serif; font-size:20px; line-height:28px; text-align:left; padding-bottom:20px;">Hello</td>
|
||||
{{end}}
|
||||
</tr>
|
||||
<tr>
|
||||
<td class="text pb20" style="color:#000000; font-family:Arial,sans-serif; font-size:14px; line-height:26px; text-align:left; padding-bottom:20px;">You or your administrator probably requested this VPN configuration. Scan the Qrcode or open the attached configuration file ({{$.Peer.GetConfigFileName}}) in the WireGuard VPN client to establish a secure VPN connection.</td>
|
||||
</tr>
|
||||
</table>
|
||||
</th>
|
||||
</td>
|
||||
</tr>
|
||||
<!-- END Button -->
|
||||
<tr>
|
||||
<td class="text" style="color:#000000; font-family:Arial,sans-serif; font-size:12px; line-height:22px; text-align:left; padding-top:20px; word-break:break-all;">If the button does not work, copy and paste the following link into your browser:<br /><a href="{{$.Link}}" target="_blank" rel="noopener noreferrer" class="link" style="color:#000000; text-decoration:underline;">{{$.Link}}</a></td>
|
||||
</tr>
|
||||
</table>
|
||||
</td>
|
||||
@@ -119,7 +119,7 @@
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
<!-- END Article / Image On The Left - Copy On The Right -->
|
||||
<!-- END Article / Copy + Download Link -->
|
||||
|
||||
<!-- Two Columns / Articles -->
|
||||
<table width="100%" border="0" cellspacing="0" cellpadding="0">
|
||||
@@ -184,4 +184,4 @@
|
||||
</tr>
|
||||
</table>
|
||||
</body>
|
||||
</html>
|
||||
</html>
|
||||
|
||||
@@ -5,8 +5,10 @@ Hello,
|
||||
{{end}}
|
||||
|
||||
You or your administrator probably requested this VPN configuration.
|
||||
Scan the attached Qrcode or open the attached configuration file ({{$.ConfigFileName}})
|
||||
in the WireGuard VPN client to establish a secure VPN connection.
|
||||
Follow the link below to download your personal WireGuard configuration file and open it
|
||||
in the WireGuard VPN client to establish a secure VPN connection:
|
||||
|
||||
{{$.Link}}
|
||||
|
||||
|
||||
|
||||
@@ -21,4 +23,4 @@ https://www.wireguard.com/install/
|
||||
|
||||
|
||||
This mail was generated by {{$.PortalName}}.
|
||||
{{$.PortalUrl}}
|
||||
{{$.PortalUrl}}
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
package domain
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Base64UrlDecode decodes a base64 url encoded string.
|
||||
// In comparison to the standard base64 encoding, the url encoding uses - instead of + and _ instead of /
|
||||
// as well as . instead of =.
|
||||
func Base64UrlDecode(in string) string {
|
||||
in = strings.ReplaceAll(in, "-", "=")
|
||||
in = strings.ReplaceAll(in, "_", "/")
|
||||
in = strings.ReplaceAll(in, ".", "+")
|
||||
|
||||
output, _ := base64.StdEncoding.DecodeString(in)
|
||||
return string(output)
|
||||
}
|
||||
|
||||
// Base64UrlEncode encodes the given input using the URL-safe base64 variant that the WireGuard Portal
|
||||
// API expects. In comparison to the standard base64 encoding, it uses . instead of +, _ instead of /
|
||||
// and - instead of =.
|
||||
func Base64UrlEncode(in string) string {
|
||||
out := base64.StdEncoding.EncodeToString([]byte(in))
|
||||
out = strings.ReplaceAll(out, "+", ".")
|
||||
out = strings.ReplaceAll(out, "/", "_")
|
||||
out = strings.ReplaceAll(out, "=", "-")
|
||||
return out
|
||||
}
|
||||
Reference in New Issue
Block a user