add header filter

This commit is contained in:
kandrusyak
2023-06-23 18:05:01 +03:00
parent eab7ec0c90
commit 7d2aaf40e4
5 changed files with 64 additions and 11 deletions

View File

@@ -23,14 +23,18 @@ func main() {
// Middleware // Middleware
e.Use(middleware.Logger()) e.Use(middleware.Logger())
e.Use(mw.FilterHeaders)
e.Use(mw.TokenFromCookies) e.Use(mw.TokenFromCookies)
e.Use(mw.JWT_MW) e.Use(mw.JWT_MW)
// Routes // Routes
e.GET("/*", checkAuth) e.GET("/*", checkAuth)
e.POST("/auth/login", users.Login)
e.GET("/health", handler.Health) e.GET("/health", handler.Health)
// Users
e.POST("/auth/login", users.Login)
e.GET("/users/me", users.GetCurrentUser)
e.Logger.Fatal(e.Start(":" + config.Env.Server.Port)) e.Logger.Fatal(e.Start(":" + config.Env.Server.Port))
} }

17
mw/filter_headers.go Normal file
View File

@@ -0,0 +1,17 @@
package mw
import (
"github.com/labstack/echo/v4"
)
var NOT_ALLOWED_HEADERS = [2]string{"UserId", "Authorization"}
func FilterHeaders(next echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error {
for _, header := range NOT_ALLOWED_HEADERS {
c.Request().Header.Set(header, "")
}
return next(c)
}
}

View File

@@ -2,22 +2,16 @@ package mw
import ( import (
"github.com/labstack/echo/v4" "github.com/labstack/echo/v4"
"net/http"
) )
func TokenFromCookies(next echo.HandlerFunc) echo.HandlerFunc { func TokenFromCookies(next echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error { return func(c echo.Context) error {
var token, err = c.Request().Cookie("accessToken") var token, err = c.Request().Cookie("accessToken")
var authToken = c.Request().Header.Get("Authorization")
if err != nil { if err != nil {
return next(c) return next(c)
} }
if authToken != "" {
return c.NoContent(http.StatusBadRequest)
}
c.Request().Header.Set("Authorization", "Bearer "+token.Value) c.Request().Header.Set("Authorization", "Bearer "+token.Value)
return next(c) return next(c)

View File

@@ -7,15 +7,39 @@ import (
"net/http" "net/http"
) )
func PostRequest(url string, model any) ([]byte, error) { func PostRequest(url string, model any, header http.Header) ([]byte, error) {
uJson, err := json.Marshal(model) uJson, err := json.Marshal(model)
if err != nil { if err != nil {
return nil, err return nil, err
} }
resp, err := http.Post(url, client := http.Client{}
"application/json", bytes.NewReader(uJson)) req, err := http.NewRequest("POST", url, bytes.NewReader(uJson))
req.Header = header
req.Header.Set("Content-Type", "application/json")
resp, err := client.Do(req)
if err != nil {
return nil, err
}
data, err := io.ReadAll(resp.Body)
if err != nil {
return nil, err
}
return data, nil
}
func GetRequest(url string, header http.Header) ([]byte, error) {
client := http.Client{}
req, err := http.NewRequest("GET", url, nil)
req.Header = header
req.Header.Set("Content-Type", "application/json")
resp, err := client.Do(req)
if err != nil { if err != nil {
return nil, err return nil, err

View File

@@ -19,7 +19,7 @@ func Login(c echo.Context) (err error) {
return return
} }
data, err := api.PostRequest(config.Env.MSHosts.UsersHost+"/login", u) data, err := api.PostRequest(config.Env.MSHosts.UsersHost+"/login", u, c.Request().Header)
u.Password = "" u.Password = ""
@@ -53,3 +53,17 @@ func Login(c echo.Context) (err error) {
return c.JSON(http.StatusOK, u) return c.JSON(http.StatusOK, u)
} }
func GetCurrentUser(c echo.Context) error {
u := new(model.User)
data, err := api.GetRequest(config.Env.MSHosts.UsersHost+"/me", c.Request().Header)
err = json.Unmarshal(data, &u)
if err != nil {
return c.NoContent(http.StatusBadGateway)
}
return c.JSON(http.StatusOK, u)
}