diff --git a/api/token.go b/api/token.go index 2672159..8aec8d8 100644 --- a/api/token.go +++ b/api/token.go @@ -14,7 +14,7 @@ import ( const defaultsecretkey = "secret" -func createJWT(secretKey string, data map[string]interface{}, ttl time.Duration) (string, error) { +func createJWT(secretKey string, data map[string]any, ttl time.Duration) (string, error) { claims := jwt.MapClaims{ "exp": time.Now().Add(ttl).Unix(), "iat": time.Now().Unix(), @@ -29,10 +29,10 @@ func createJWT(secretKey string, data map[string]interface{}, ttl time.Duration) return signedToken, nil } -func VerifyJWT(tokenString string, secretKey string) (map[string]interface{}, error) { +func VerifyJWT(tokenString string, secretKey string) (map[string]any, error) { token, e := jwt.Parse( tokenString, - func(token *jwt.Token) (interface{}, error) { + func(token *jwt.Token) (any, error) { if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok { return nil, fmt.Errorf("unexpexted signing method: %v", token.Header["alg"]) } @@ -43,7 +43,7 @@ func VerifyJWT(tokenString string, secretKey string) (map[string]interface{}, er return nil, fmt.Errorf("token validation failed: %w", e) } if claims, ok := token.Claims.(jwt.MapClaims); ok && token.Valid { - result := make(map[string]interface{}) + result := make(map[string]any) for key, value := range claims { if key == "exp" || key == "iat" || key == "nbf" { continue @@ -74,7 +74,7 @@ func Token(w http.ResponseWriter, req *http.Request) { IsSuperuser: dbuser.IsSuperuser, } responseUserMarshal, _ := json.Marshal(responseUser) - responseUserUnmarshal := make(map[string]interface{}) + responseUserUnmarshal := make(map[string]any) json.Unmarshal(responseUserMarshal, &responseUserUnmarshal) token, _ := createJWT(defaultsecretkey, responseUserUnmarshal, time.Duration(24*time.Hour)) fmt.Fprintf(w, "{\"token\":\"%s\"}", token) diff --git a/api/users.go b/api/users.go index 1f91ea7..a70361b 100644 --- a/api/users.go +++ b/api/users.go @@ -6,6 +6,8 @@ import ( "form/database" "form/model" "net/http" + + "github.com/google/uuid" ) func GetAllUsers(w http.ResponseWriter, req *http.Request) { @@ -65,25 +67,156 @@ func GetUser(w http.ResponseWriter, req *http.Request) { func ChangePassword(w http.ResponseWriter, req *http.Request) { w.Header().Set("Content-Type", "application/json") - + var request model.ChangePasswordRequest + decoder := json.NewDecoder(req.Body) + decoder.Decode(&request) + claims, e := VerifyJWT(request.Token, defaultsecretkey) + if e != nil { + fmt.Fprintf(w, "{\"status\": 403}") + return + } + user, e := database.GetUserByUsername(fmt.Sprintf("%v", claims["username"])) + if e != nil { + fmt.Fprintf(w, "{\"error\":\"%v\"}", e) + return + } + if user.Password == request.OldPassword { + database.Db.Model(&user).Update("Password", request.NewPassword) + fmt.Fprint(w, "{\"status\":200}") + return + } + fmt.Fprintf(w, "{\"status\":403}") } func DelUser(w http.ResponseWriter, req *http.Request) { w.Header().Set("Content-Type", "application/json") - + var request model.DelUserRequest + decoder := json.NewDecoder(req.Body) + decoder.Decode(&request) + claims, e := VerifyJWT(request.Token, defaultsecretkey) + if e != nil { + fmt.Fprintf(w, "{\"status\": 403}") + return + } + if claims["is_superuser"] == true || request.UserId != "" { + user, e := database.GetUserById(request.UserId) + if e != nil { + fmt.Fprintf(w, "{\"error\":\"%v\"}", e) + return + } + e = database.DeleteUser(user.Id) + if e != nil { + fmt.Fprintf(w, "{\"error\":\"%v\"}", e) + return + } + fmt.Fprintf(w, "{\"status\":200}") + return + } + e = database.DeleteUser(fmt.Sprintf("%v", claims["id"])) + if e != nil { + fmt.Fprintf(w, "{\"error\":\"%v\"}", e) + return + } + fmt.Fprint(w, "{\"status\":200}") } func AddUser(w http.ResponseWriter, req *http.Request) { w.Header().Set("Content-Type", "application/json") - + var request model.AddUserRequest + decoder := json.NewDecoder(req.Body) + decoder.Decode(&request) + claims, e := VerifyJWT(request.Token, defaultsecretkey) + if e != nil { + fmt.Fprintf(w, "{\"status\": 403}") + return + } + if claims["is_superuser"] == true { + id, err := uuid.NewRandom() + if err != nil { + fmt.Fprintf(w, "{\"error\":\"%v\"}", err) + return + } + _, e := database.GetUserByUsername(request.Username) + if e != nil { + e = database.CreateUser(&model.User{ + Id: id.String(), + Username: request.Username, + Email: request.Email, + Password: request.Password, + Group: request.Group, + IsSuperuser: request.IsSuperuser, + }) + if e != nil { + fmt.Fprintf(w, "{\"error\":\"%v\"}", e) + return + } + fmt.Fprintf(w, "{\"status\": 200}") + return + } + fmt.Fprintf(w, "{\"error\":\"username already exists\"}") + return + } + fmt.Fprintf(w, "{\"status\":403}") } func RegisterUser(w http.ResponseWriter, req *http.Request) { w.Header().Set("Content-Type", "application/json") - + var request model.AddUserRequest + decoder := json.NewDecoder(req.Body) + decoder.Decode(&request) + if request.Username == "" { + fmt.Fprintf(w, "{\"error\":\"need more data\"}") + return + } + id, err := uuid.NewRandom() + if err != nil { + fmt.Fprintf(w, "{\"error\":\"%v\"}", err) + return + } + group := "users" + if request.Group != "" { + group = request.Group + } + _, e := database.GetUserByUsername(request.Username) + if e != nil { + e = database.CreateUser(&model.User{ + Id: id.String(), + Username: request.Username, + Email: request.Email, + Password: request.Password, + Group: group, + IsSuperuser: request.IsSuperuser, + }) + if e != nil { + fmt.Fprintf(w, "{\"error\":\"%v\"}", e) + return + } + fmt.Fprintf(w, "{\"status\": 200}") + return + } + fmt.Fprintf(w, "{\"error\":\"username already exists\"}") } func ChangeUser(w http.ResponseWriter, req *http.Request) { w.Header().Set("Content-Type", "application/json") - + var request model.ChangeUserRequest + decoder := json.NewDecoder(req.Body) + decoder.Decode(&request) + claims, e := VerifyJWT(request.Token, defaultsecretkey) + if e != nil { + fmt.Fprintf(w, "{\"status\": 403}") + return + } + user, e := database.GetUserByUsername(fmt.Sprintf("%v", claims["username"])) + if e != nil { + fmt.Fprintf(w, "{\"error\":\"%v\"}", e) + return + } + database.Db.Model(&user).Updates(map[string]any{ + "Username": request.Username, + "Email": request.Email, + "Group": request.Group, + "IsSuperuser": request.IsSuperuser, + }) + fmt.Fprint(w, "{\"status\":200}") } diff --git a/auth_service.config b/auth_service.config new file mode 100644 index 0000000..0ac9aa0 --- /dev/null +++ b/auth_service.config @@ -0,0 +1,3 @@ +USERS_FILE=users.json +USERS_DB=users.db +ENABLE_SELF_REGISTER=false \ No newline at end of file diff --git a/database/db.go b/database/db.go index 217cd58..8ff1f86 100644 --- a/database/db.go +++ b/database/db.go @@ -3,16 +3,17 @@ package database import ( "fmt" "form/model" + "os" "gorm.io/driver/sqlite" "gorm.io/gorm" ) -var db *gorm.DB +var Db *gorm.DB func Init() error { var err error - db, err = gorm.Open(sqlite.Open("users.db"), &gorm.Config{}) + Db, err = gorm.Open(sqlite.Open(os.Getenv("USERS_DB")), &gorm.Config{}) if err != nil { fmt.Println("failed to connect database") return err @@ -21,5 +22,5 @@ func Init() error { } func Migrate() { - db.AutoMigrate(&model.User{}) + Db.AutoMigrate(&model.User{}) } diff --git a/database/users.go b/database/users.go index 20e839b..5089954 100644 --- a/database/users.go +++ b/database/users.go @@ -12,7 +12,7 @@ import ( ) func CreateUser(user *model.User) error { - e := db.Where(model.User{ + e := Db.Where(model.User{ Username: user.Username, }).FirstOrCreate(&user).Error if e != nil { @@ -23,7 +23,7 @@ func CreateUser(user *model.User) error { func GetUserById(id string) (*model.User, error) { var user model.User - e := db.First(&user, "id = ?", id).Error + e := Db.First(&user, "id = ?", id).Error if e != nil { if errors.Is(e, gorm.ErrRecordNotFound) { return nil, fmt.Errorf("user not found") @@ -35,7 +35,7 @@ func GetUserById(id string) (*model.User, error) { func GetUserByUsername(username string) (*model.User, error) { var user model.User - e := db.First(&user, "username = ?", username).Error + e := Db.First(&user, "username = ?", username).Error if e != nil { if errors.Is(e, gorm.ErrRecordNotFound) { return nil, fmt.Errorf("user not found") @@ -47,14 +47,14 @@ func GetUserByUsername(username string) (*model.User, error) { func DeleteUser(id string) error { var user model.User - e := db.First(&user, "id = ?", id).Error + e := Db.First(&user, "id = ?", id).Error if e != nil { if errors.Is(e, gorm.ErrRecordNotFound) { return fmt.Errorf("user not found") } return e } - e = db.Delete(&user).Error + e = Db.Unscoped().Delete(&user).Error if e != nil { if errors.Is(e, gorm.ErrRecordNotFound) { return fmt.Errorf("user not found") @@ -77,7 +77,7 @@ func GetUsersFromFile(filename string) []model.User { func GetAllUsers() ([]model.User, error) { var users []model.User - e := db.Find(&users).Error + e := Db.Find(&users).Error if e != nil { return nil, e } diff --git a/go.mod b/go.mod index 0c62eab..0cd5223 100644 --- a/go.mod +++ b/go.mod @@ -4,6 +4,8 @@ go 1.24.3 require ( github.com/golang-jwt/jwt/v5 v5.2.2 + github.com/google/uuid v1.6.0 + github.com/joho/godotenv v1.5.1 gorm.io/driver/sqlite v1.5.7 gorm.io/gorm v1.30.0 ) diff --git a/main.go b/main.go index f715110..c58705a 100644 --- a/main.go +++ b/main.go @@ -3,28 +3,40 @@ package main import ( "form/api" "form/database" + "log" "net/http" + "os" + + "github.com/joho/godotenv" ) func main() { + err := godotenv.Load("auth_service.config") + if err != nil { + log.Fatal("Error loading .env file") + } database.Init() database.Migrate() - users := database.GetUsersFromFile("users.json") - for i := range users { - database.CreateUser(&users[i]) + if os.Getenv("USERS_FILE") != "" { + users := database.GetUsersFromFile(os.Getenv("USERS_FILE")) + for i := range users { + database.CreateUser(&users[i]) + } } http.HandleFunc("/api/token", api.Token) http.HandleFunc("/api/token/verify", api.VerifyToken) http.HandleFunc("/api/access", api.VerifyToken) - http.HandleFunc("/api/register", api.RegisterUser) + if os.Getenv("ENABLE_SELF_REGISTER") == "true" { + http.HandleFunc("/api/register", api.RegisterUser) + } http.HandleFunc("/api/users", api.GetAllUsers) http.HandleFunc("/api/users/user", api.GetUser) http.HandleFunc("/api/users/password", api.ChangePassword) http.HandleFunc("/api/users/change", api.ChangeUser) http.HandleFunc("/api/users/del", api.DelUser) - http.HandleFunc("/api/user/add", api.AddUser) + http.HandleFunc("/api/users/add", api.AddUser) http.ListenAndServe(":8090", nil) } diff --git a/model/user.go b/model/user.go index ddba06c..65ad22d 100644 --- a/model/user.go +++ b/model/user.go @@ -32,3 +32,39 @@ type Token struct { type GetAllUsersRequest struct { Token string `json:"token"` } + +type ChangePasswordRequest struct { + Token string `json:"token"` + OldPassword string `json:"old_password"` + NewPassword string `json:"new_password"` +} + +type ChangeUserRequest struct { + Token string `json:"token"` + Username string `json:"username"` + Email string `json:"email"` + Group string `json:"group"` + IsSuperuser bool `json:"is_superuser"` +} + +type DelUserRequest struct { + Token string `json:"token"` + UserId string `json:"user_id,omitempty"` +} + +type AddUserRequest struct { + Token string `json:"token"` + Username string `json:"username"` + Email string `json:"email"` + Password string `json:"password"` + Group string `json:"group"` + IsSuperuser bool `json:"is_superuser"` +} + +type RegusterUserRequest struct { + Username string `json:"username"` + Email string `json:"email"` + Password string `json:"password"` + Group string `json:"group"` + IsSuperuser bool `json:"is_superuser"` +}