mgmt_api.go 7.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269
  1. package http
  2. import (
  3. "encoding/base64"
  4. "encoding/json"
  5. "errors"
  6. "fmt"
  7. "log/slog"
  8. "net"
  9. "net/http"
  10. "os"
  11. "strings"
  12. "github.com/google/uuid"
  13. "github.com/mk6i/retro-aim-server/config"
  14. "github.com/mk6i/retro-aim-server/state"
  15. "github.com/mk6i/retro-aim-server/wire"
  16. )
  17. type userWithPassword struct {
  18. state.User
  19. Password string `json:"password,omitempty"`
  20. }
  21. type userSession struct {
  22. ScreenName string `json:"screen_name"`
  23. }
  24. type onlineUsers struct {
  25. Count int `json:"count"`
  26. Sessions []userSession `json:"sessions"`
  27. }
  28. type UserManager interface {
  29. AllUsers() ([]state.User, error)
  30. InsertUser(u state.User) error
  31. SetUserPassword(u state.User) error
  32. User(screenName string) (*state.User, error)
  33. }
  34. type SessionRetriever interface {
  35. AllSessions() []*state.Session
  36. }
  37. func StartManagementAPI(cfg config.Config, userManager UserManager, sessionRetriever SessionRetriever, logger *slog.Logger) {
  38. mux := http.NewServeMux()
  39. newUser := func() state.User {
  40. return state.User{AuthKey: uuid.New().String()}
  41. }
  42. mux.HandleFunc("/user", func(w http.ResponseWriter, r *http.Request) {
  43. userHandler(w, r, userManager, newUser, logger)
  44. })
  45. mux.HandleFunc("/user/password", func(w http.ResponseWriter, r *http.Request) {
  46. userPasswordHandler(w, r, userManager, newUser, logger)
  47. })
  48. mux.HandleFunc("/user/login", func(w http.ResponseWriter, r *http.Request) {
  49. loginHandler(w, r, userManager, logger)
  50. })
  51. mux.HandleFunc("/session", func(w http.ResponseWriter, r *http.Request) {
  52. sessionHandler(w, r, sessionRetriever)
  53. })
  54. addr := net.JoinHostPort(cfg.ApiHost, cfg.ApiPort)
  55. logger.Info("starting management API server", "addr", addr)
  56. if err := http.ListenAndServe(addr, mux); err != nil {
  57. logger.Error("unable to bind management API address address", "err", err.Error())
  58. os.Exit(1)
  59. }
  60. }
  61. func userHandler(
  62. w http.ResponseWriter,
  63. r *http.Request,
  64. userManager UserManager,
  65. newUser func() state.User,
  66. logger *slog.Logger,
  67. ) {
  68. switch r.Method {
  69. case http.MethodGet:
  70. getUserHandler(w, r, userManager, logger)
  71. case http.MethodPost:
  72. postUserHandler(w, r, userManager, newUser, logger)
  73. default:
  74. http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
  75. }
  76. }
  77. func userPasswordHandler(
  78. w http.ResponseWriter,
  79. r *http.Request,
  80. userManager UserManager,
  81. userFactory func() state.User,
  82. logger *slog.Logger,
  83. ) {
  84. switch r.Method {
  85. case http.MethodPut:
  86. putUserPasswordHandler(w, r, userManager, userFactory, logger)
  87. default:
  88. http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
  89. }
  90. }
  91. // putUserPasswordHandler handles the PUT /user/password endpoint.
  92. func putUserPasswordHandler(
  93. w http.ResponseWriter,
  94. r *http.Request,
  95. userManager UserManager,
  96. newUser func() state.User,
  97. logger *slog.Logger,
  98. ) {
  99. user := userWithPassword{
  100. User: newUser(),
  101. }
  102. if err := json.NewDecoder(r.Body).Decode(&user); err != nil {
  103. http.Error(w, "malformed input", http.StatusBadRequest)
  104. return
  105. }
  106. if err := user.HashPassword(user.Password); err != nil {
  107. logger.Error("error hashing user password in PUT /user/password", "err", err.Error())
  108. http.Error(w, "internal server error", http.StatusInternalServerError)
  109. return
  110. }
  111. if err := userManager.SetUserPassword(user.User); err != nil {
  112. switch {
  113. case errors.Is(err, state.ErrNoUser):
  114. http.Error(w, "user does not exist", http.StatusNotFound)
  115. return
  116. case err != nil:
  117. logger.Error("error updating user password PUT /user/password", "err", err.Error())
  118. http.Error(w, "internal server error", http.StatusInternalServerError)
  119. return
  120. }
  121. }
  122. w.WriteHeader(http.StatusNoContent)
  123. }
  124. // sessionHandler handles GET /session
  125. func sessionHandler(w http.ResponseWriter, r *http.Request, sessionRetriever SessionRetriever) {
  126. w.Header().Set("Content-Type", "application/json")
  127. if r.Method != http.MethodGet {
  128. http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
  129. return
  130. }
  131. allUsers := sessionRetriever.AllSessions()
  132. ou := onlineUsers{
  133. Count: len(allUsers),
  134. Sessions: make([]userSession, 0),
  135. }
  136. for _, s := range allUsers {
  137. ou.Sessions = append(ou.Sessions, userSession{
  138. ScreenName: s.ScreenName(),
  139. })
  140. }
  141. if err := json.NewEncoder(w).Encode(ou); err != nil {
  142. http.Error(w, err.Error(), http.StatusInternalServerError)
  143. return
  144. }
  145. }
  146. // getUserHandler handles the GET /user endpoint.
  147. func getUserHandler(w http.ResponseWriter, _ *http.Request, userManager UserManager, logger *slog.Logger) {
  148. w.Header().Set("Content-Type", "application/json")
  149. users, err := userManager.AllUsers()
  150. if err != nil {
  151. logger.Error("error in GET /user", "err", err.Error())
  152. http.Error(w, "internal server error", http.StatusInternalServerError)
  153. return
  154. }
  155. if err := json.NewEncoder(w).Encode(users); err != nil {
  156. http.Error(w, err.Error(), http.StatusInternalServerError)
  157. return
  158. }
  159. }
  160. // postUserHandler handles the POST /user endpoint.
  161. func postUserHandler(
  162. w http.ResponseWriter,
  163. r *http.Request,
  164. userManager UserManager,
  165. newUser func() state.User,
  166. logger *slog.Logger,
  167. ) {
  168. user := userWithPassword{
  169. User: newUser(),
  170. }
  171. if err := json.NewDecoder(r.Body).Decode(&user); err != nil {
  172. http.Error(w, "malformed input", http.StatusBadRequest)
  173. return
  174. }
  175. if err := user.HashPassword(user.Password); err != nil {
  176. logger.Error("error hashing user password in POST /user", "err", err.Error())
  177. http.Error(w, "internal server error", http.StatusInternalServerError)
  178. return
  179. }
  180. err := userManager.InsertUser(user.User)
  181. switch {
  182. case errors.Is(err, state.ErrDupUser):
  183. http.Error(w, "user already exists", http.StatusConflict)
  184. return
  185. case err != nil:
  186. logger.Error("error inserting user POST /user", "err", err.Error())
  187. http.Error(w, "internal server error", http.StatusInternalServerError)
  188. return
  189. }
  190. w.WriteHeader(http.StatusCreated)
  191. fmt.Fprintln(w, "User account created successfully.")
  192. }
  193. // loginHandler is a temporary endpoint for validating user credentials for
  194. // chivanet. do not rely on this endpoint, as it will be eventually removed.
  195. func loginHandler(w http.ResponseWriter, r *http.Request, userManager UserManager, logger *slog.Logger) {
  196. authHeader := r.Header.Get("Authorization")
  197. if authHeader == "" {
  198. // No authentication header found
  199. w.WriteHeader(http.StatusUnauthorized)
  200. w.Header().Set("WWW-Authenticate", `Basic realm="User Login"`)
  201. w.Write([]byte("401 Unauthorized\n"))
  202. return
  203. }
  204. auth := strings.SplitN(authHeader, " ", 2)
  205. if len(auth) != 2 || auth[0] != "Basic" {
  206. w.WriteHeader(http.StatusUnauthorized)
  207. w.Write([]byte("401 Unauthorized: Missing Basic prefix\n"))
  208. return
  209. }
  210. payload, err := base64.StdEncoding.DecodeString(auth[1])
  211. if err != nil {
  212. w.WriteHeader(http.StatusUnauthorized)
  213. w.Write([]byte("401 Unauthorized: Invalid Base64 Encoding\n"))
  214. return
  215. }
  216. pair := strings.SplitN(string(payload), ":", 2)
  217. if len(pair) != 2 {
  218. w.WriteHeader(http.StatusUnauthorized)
  219. w.Write([]byte("401 Unauthorized: Invalid Authentication Token\n"))
  220. return
  221. }
  222. username, password := pair[0], pair[1]
  223. user, err := userManager.User(username)
  224. if err != nil {
  225. w.WriteHeader(http.StatusInternalServerError)
  226. w.Write([]byte("500 InternalServerError\n"))
  227. logger.Error("error getting user", "err", err.Error())
  228. return
  229. }
  230. if user == nil || !user.ValidateHash(wire.StrongMD5PasswordHash(password, user.AuthKey)) {
  231. w.WriteHeader(http.StatusUnauthorized)
  232. w.Write([]byte("401 Unauthorized: Invalid Credentials\n"))
  233. return
  234. }
  235. // Successfully authenticated
  236. w.WriteHeader(http.StatusOK)
  237. w.Write([]byte("200 OK: Successfully Authenticated\n"))
  238. }