Skip to content
Merged
22 changes: 11 additions & 11 deletions internals/config/loader.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ import (
"strings"

"github.com/codeshelldev/gotl/pkg/configutils"
log "github.com/codeshelldev/gotl/pkg/logger"
"github.com/codeshelldev/gotl/pkg/logger"
"github.com/codeshelldev/gotl/pkg/stringutils"
"github.com/codeshelldev/secured-signal-api/internals/config/structure"

Expand Down Expand Up @@ -63,13 +63,13 @@ func Load() {

InitTokens()

log.Info("Finished Loading Configuration")
logger.Info("Finished Loading Configuration")
}

func Log() {
log.Dev("Loaded Config:", mainConf.Layer.Get(""))
log.Dev("Loaded Token Configs:", tokenConf.Layer.Get(""))
log.Dev("Parsed Configs: ", ENV)
logger.Dev("Loaded Config:", mainConf.Layer.Get(""))
logger.Dev("Loaded Token Configs:", tokenConf.Layer.Get(""))
logger.Dev("Parsed Configs: ", ENV)
}

func Clear() {
Expand Down Expand Up @@ -102,7 +102,7 @@ func Normalize(id string, config *configutils.Config, path string, structure any
old, ok := data.(map[string]any)

if !ok {
log.Warn("Could not load `"+path+"`")
logger.Warn("Could not load `"+path+"`")
return
}

Expand All @@ -124,7 +124,7 @@ func Normalize(id string, config *configutils.Config, path string, structure any

func InitReload() {
reload := func(path string) {
log.Debug(path, " changed, reloading...")
logger.Debug(path, " changed, reloading...")
Load()
Log()
}
Expand All @@ -145,16 +145,16 @@ func InitConfig() {
}

func LoadDefaults() {
log.Debug("Loading defaults ", ENV.DEFAULTS_PATH)
logger.Debug("Loading defaults ", ENV.DEFAULTS_PATH)
_, err := defaultsConf.LoadFile(ENV.DEFAULTS_PATH, yaml.Parser())

if err != nil {
log.Warn("Could not Load Defaults", ENV.DEFAULTS_PATH)
logger.Warn("Could not Load Defaults", ENV.DEFAULTS_PATH)
}
}

func LoadConfig() {
log.Debug("Loading Config ", ENV.CONFIG_PATH)
logger.Debug("Loading Config ", ENV.CONFIG_PATH)
_, err := userConf.LoadFile(ENV.CONFIG_PATH, yaml.Parser())

if err != nil {
Expand All @@ -166,7 +166,7 @@ func LoadConfig() {
return
}

log.Error("Could not Load Config ", ENV.CONFIG_PATH, ": ", err.Error())
logger.Error("Could not Load Config ", ENV.CONFIG_PATH, ": ", err.Error())
}
}

Expand Down
12 changes: 6 additions & 6 deletions internals/config/tokens.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,18 +4,18 @@ import (
"strconv"

"github.com/codeshelldev/gotl/pkg/configutils"
log "github.com/codeshelldev/gotl/pkg/logger"
"github.com/codeshelldev/gotl/pkg/logger"
"github.com/codeshelldev/secured-signal-api/internals/config/structure"
"github.com/knadh/koanf/parsers/yaml"
)

func LoadTokens() {
log.Debug("Loading Configs in ", ENV.TOKENS_DIR)
logger.Debug("Loading Configs in ", ENV.TOKENS_DIR)

err := tokenConf.LoadDir("tokenconfigs", ENV.TOKENS_DIR, ".yml", yaml.Parser())

if err != nil {
log.Error("Could not Load Configs in ", ENV.TOKENS_DIR, ": ", err.Error())
logger.Error("Could not Load Configs in ", ENV.TOKENS_DIR, ": ", err.Error())
}

tokenConf.TemplateConfig()
Expand Down Expand Up @@ -57,9 +57,9 @@ func InitTokens() {
}

if len(apiTokens) <= 0 {
log.Warn("No API Tokens provided this is NOT recommended")
logger.Warn("No API Tokens provided this is NOT recommended")

log.Info("Disabling Security Features due to incomplete Congfiguration")
logger.Info("Disabling Security Features due to incomplete Congfiguration")

ENV.INSECURE = true

Expand All @@ -69,7 +69,7 @@ func InitTokens() {
}

if len(apiTokens) > 0 {
log.Debug("Registered " + strconv.Itoa(len(apiTokens)) + " Tokens")
logger.Debug("Registered " + strconv.Itoa(len(apiTokens)) + " Tokens")
}
}

Expand Down
97 changes: 50 additions & 47 deletions internals/proxy/middlewares/auth.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,6 @@ import (
"strings"

"github.com/codeshelldev/gotl/pkg/logger"
log "github.com/codeshelldev/gotl/pkg/logger"
"github.com/codeshelldev/gotl/pkg/request"
"github.com/codeshelldev/secured-signal-api/internals/config"
)
Expand All @@ -21,126 +20,128 @@ var Auth Middleware = Middleware{
Use: authHandler,
}

const tokenKey contextKey = "token"

type AuthMethod struct {
Name string
Authenticate func(w http.ResponseWriter, req *http.Request, tokens []string) (bool, error)
Authenticate func(w http.ResponseWriter, req *http.Request, tokens []string) (string, error)
}

var BearerAuth = AuthMethod {
Name: "Bearer",
Authenticate: func(w http.ResponseWriter, req *http.Request, tokens []string) (bool, error) {
Authenticate: func(w http.ResponseWriter, req *http.Request, tokens []string) (string, error) {
header := req.Header.Get("Authorization")

headerParts := strings.SplitN(header, " ", 2)

if len(headerParts) != 2 {
return false, nil
return "", nil
}

if strings.ToLower(headerParts[0]) == "bearer" {
if isValidToken(tokens, headerParts[1]) {
return true, nil
return headerParts[1], nil
}

return false, errors.New("invalid Bearer token")
return "", errors.New("invalid Bearer token")
}

return false, nil
return "", nil
},
}

var BasicAuth = AuthMethod {
Name: "Basic",
Authenticate: func(w http.ResponseWriter, req *http.Request, tokens []string) (bool, error) {
Authenticate: func(w http.ResponseWriter, req *http.Request, tokens []string) (string, error) {
header := req.Header.Get("Authorization")

if strings.TrimSpace(header) == "" {
return false, nil
return "", nil
}

headerParts := strings.SplitN(header, " ", 2)

if len(headerParts) != 2 {
return false, nil
return "", nil
}

if strings.ToLower(headerParts[0]) == "basic" {
base64Bytes, err := base64.StdEncoding.DecodeString(headerParts[1])

if err != nil {
log.Error("Could not decode Basic auth payload: ", err.Error())
return false, errors.New("invalid base64 in Basic auth")
logger.Error("Could not decode Basic auth payload: ", err.Error())
return "", errors.New("invalid base64 in Basic auth")
}

parts := strings.SplitN(string(base64Bytes), ":", 2)

if len(parts) != 2 {
return false, errors.New("Basic auth must be user:password")
return "", errors.New("Basic auth must be user:password")
}

user, password := parts[0], parts[1]

if strings.ToLower(user) == "api" && isValidToken(tokens, password) {
return true, nil
return password, nil
}

return false, errors.New("invalid user:password")
return "", errors.New("invalid user:password")
}

return false, nil
return "", nil
},
}

var BodyAuth = AuthMethod {
Name: "Body",
Authenticate: func(w http.ResponseWriter, req *http.Request, tokens []string) (bool, error) {
Authenticate: func(w http.ResponseWriter, req *http.Request, tokens []string) (string, error) {
const authField = "auth"

body, err := request.GetReqBody(req)

if err != nil {
return false, nil
return "", nil
}

body.Write(req)

if body.Empty {
return false, nil
return "", nil
}

value, exists := body.Data[authField]

if !exists {
return false, nil
return "", nil
}

auth, ok := value.(string)

if !ok {
return false, nil
return "", nil
}

if isValidToken(tokens, auth) {
delete(body.Data, authField)

body.Write(req)

return true, nil
return auth, nil
}

return false, errors.New("invalid Body token")
return "", errors.New("invalid Body token")
},
}

var QueryAuth = AuthMethod {
Name: "Query",
Authenticate: func(w http.ResponseWriter, req *http.Request, tokens []string) (bool, error) {
Authenticate: func(w http.ResponseWriter, req *http.Request, tokens []string) (string, error) {
const authQuery = "@authorization"

auth := req.URL.Query().Get(authQuery)

if strings.TrimSpace(auth) == "" {
return false, nil
return "", nil
}

if isValidToken(tokens, auth) {
Expand All @@ -150,39 +151,39 @@ var QueryAuth = AuthMethod {

req.URL.RawQuery = query.Encode()

return true, nil
return auth, nil
}

return false, errors.New("invalid Query token")
return "", errors.New("invalid Query token")
},
}

var PathAuth = AuthMethod {
Name: "Path",
Authenticate: func(w http.ResponseWriter, req *http.Request, tokens []string) (bool, error) {
Authenticate: func(w http.ResponseWriter, req *http.Request, tokens []string) (string, error) {
parts := strings.Split(req.URL.Path, "/")

if len(parts) == 0 {
return false, nil
return "", nil
}

unescaped, err := url.PathUnescape(parts[1])

if err != nil {
return false, nil
return "", nil
}

auth, exists := strings.CutPrefix(unescaped, "auth=")

if !exists {
return false, nil
return "", nil
}

if isValidToken(tokens, auth) {
return true, nil
return auth, nil
}

return false, errors.New("invalid Path token")
return "", errors.New("invalid Path token")
},
}

Expand All @@ -207,24 +208,26 @@ func authHandler(next http.Handler) http.Handler {
return
}

var authToken string

success, _ := authChain.Eval(w, req, tokens)
token, _ := authChain.Eval(w, req, tokens)

if !success {
logger.Warn("User failed to provide auth")
w.Header().Set("WWW-Authenticate", "Basic realm=\"Login Required\", Bearer realm=\"Access Token Required\"")
http.Error(w, "Unauthorized", http.StatusUnauthorized)
if token == "" {
onUnauthorized(w)
return
}

ctx := context.WithValue(req.Context(), tokenKey, authToken)
ctx := context.WithValue(req.Context(), tokenKey, token)
req = req.WithContext(ctx)

next.ServeHTTP(w, req)
})
}

func onUnauthorized(w http.ResponseWriter) {
w.Header().Set("WWW-Authenticate", "Basic realm=\"Login Required\", Bearer realm=\"Access Token Required\"")

http.Error(w, "Unauthorized", http.StatusUnauthorized)
}

func isValidToken(tokens []string, match string) bool {
return slices.Contains(tokens, match)
}
Expand All @@ -245,21 +248,21 @@ func (chain *AuthChain) Use(method AuthMethod) *AuthChain {
return chain
}

func (chain *AuthChain) Eval(w http.ResponseWriter, req *http.Request, tokens []string) (bool, error) {
func (chain *AuthChain) Eval(w http.ResponseWriter, req *http.Request, tokens []string) (string, error) {
var err error
var success bool
var token string

for _, method := range chain.methods {
success, err = method.Authenticate(w, req, tokens)
token, err = method.Authenticate(w, req, tokens)

if err != nil {
logger.Warn("User failed ", method.Name, " auth: ", err.Error())
}

if success {
return success, nil
if token != "" {
return token, nil
}
}

return false, err
return "", err
}
2 changes: 0 additions & 2 deletions internals/proxy/middlewares/common.go
Original file line number Diff line number Diff line change
Expand Up @@ -13,8 +13,6 @@ type Context struct {

type contextKey string

const tokenKey contextKey = "token"

func getConfigByReq(req *http.Request) *structure.CONFIG {
token := req.Context().Value(tokenKey).(string)

Expand Down
Loading