111 lines
4.2 KiB
Go
111 lines
4.2 KiB
Go
|
|
package server
|
||
|
|
|
||
|
|
import (
|
||
|
|
"errors"
|
||
|
|
"log/slog"
|
||
|
|
"net/http"
|
||
|
|
"net/url"
|
||
|
|
"strings"
|
||
|
|
|
||
|
|
"github.com/go-chi/chi/v5"
|
||
|
|
|
||
|
|
"glchat/internal/auth"
|
||
|
|
"glchat/internal/httpx"
|
||
|
|
)
|
||
|
|
|
||
|
|
// registerOAuthRoutes вешает ручки входа через внешние провайдеры (AGENT.md 7.1).
|
||
|
|
// `start` и `callback` — публичные (это вход), лимит по IP общий с другими
|
||
|
|
// ручками аутентификации.
|
||
|
|
func (s *Server) registerOAuthRoutes(router chi.Router) {
|
||
|
|
oauthLimit := s.oauthLimiter.Middleware(func(r *http.Request) string {
|
||
|
|
return httpx.ClientIP(r, nil)
|
||
|
|
})
|
||
|
|
router.Get("/auth/oauth/providers", s.handleOAuthProviders)
|
||
|
|
router.With(oauthLimit).Get("/auth/oauth/{provider}/start", s.handleOAuthStart)
|
||
|
|
router.With(oauthLimit).Get("/auth/oauth/{provider}/callback", s.handleOAuthCallback)
|
||
|
|
}
|
||
|
|
|
||
|
|
// handleOAuthProviders отдаёт список включённых провайдеров: клиент рисует по
|
||
|
|
// нему кнопки входа, а закрытые провайдеры не показывает вовсе.
|
||
|
|
func (s *Server) handleOAuthProviders(w http.ResponseWriter, r *http.Request) {
|
||
|
|
if s.auth == nil {
|
||
|
|
writeJSON(w, map[string]any{"providers": []any{}, "redirect_url": ""})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
writeJSON(w, map[string]any{
|
||
|
|
"providers": s.auth.OAuthProviders(),
|
||
|
|
"redirect_url": s.auth.OAuthRedirectURL(),
|
||
|
|
})
|
||
|
|
}
|
||
|
|
|
||
|
|
// handleOAuthStart уводит браузер на страницу согласия провайдера.
|
||
|
|
// Если провайдер не настроен — отвечаем понятной ошибкой, а не редиректом.
|
||
|
|
func (s *Server) handleOAuthStart(w http.ResponseWriter, r *http.Request) {
|
||
|
|
if s.auth == nil {
|
||
|
|
writeAPIError(w, auth.ErrSessionExpired)
|
||
|
|
return
|
||
|
|
}
|
||
|
|
provider := chi.URLParam(r, "provider")
|
||
|
|
redirect := r.URL.Query().Get("redirect")
|
||
|
|
authorizeURL, err := s.auth.OAuthAuthorizeURL(provider, redirect)
|
||
|
|
if err != nil {
|
||
|
|
writeAPIError(w, err)
|
||
|
|
return
|
||
|
|
}
|
||
|
|
//nolint:gosec // адрес берётся из каталога провайдеров, пользователь его не задаёт
|
||
|
|
http.Redirect(w, r, authorizeURL, http.StatusFound)
|
||
|
|
}
|
||
|
|
|
||
|
|
// handleOAuthCallback принимает код, находит или создаёт аккаунт и возвращает
|
||
|
|
// пользователя в клиент. Ошибки уходят в `?oauth_error=<код>`: браузер видит
|
||
|
|
// понятное сообщение на странице входа, а не JSON.
|
||
|
|
func (s *Server) handleOAuthCallback(w http.ResponseWriter, r *http.Request) {
|
||
|
|
if s.auth == nil {
|
||
|
|
writeAPIError(w, auth.ErrSessionExpired)
|
||
|
|
return
|
||
|
|
}
|
||
|
|
provider := chi.URLParam(r, "provider")
|
||
|
|
query := r.URL.Query()
|
||
|
|
_, token, session, redirect, err := s.auth.OAuthCallback(r.Context(), provider,
|
||
|
|
query.Get("code"), query.Get("state"),
|
||
|
|
httpx.ClientIPFromContext(r.Context()), httpx.UserAgentFromContext(r.Context()))
|
||
|
|
if err != nil {
|
||
|
|
code := oauthErrorCode(err)
|
||
|
|
s.logger.InfoContext(r.Context(), "вход через OAuth не удался",
|
||
|
|
slog.String("provider", provider), slog.String("code", code))
|
||
|
|
http.Redirect(w, r, "/login?oauth_error="+url.QueryEscape(code), http.StatusFound)
|
||
|
|
return
|
||
|
|
}
|
||
|
|
http.SetCookie(w, s.sessionCookie(token, session.ExpiresAt))
|
||
|
|
// Признак в адресе нужен клиенту, чтобы перечитать профиль после возврата.
|
||
|
|
target := redirect
|
||
|
|
if target == "" || target == "/" {
|
||
|
|
target = "/app"
|
||
|
|
}
|
||
|
|
separator := "?"
|
||
|
|
if strings.Contains(target, "?") {
|
||
|
|
separator = "&"
|
||
|
|
}
|
||
|
|
//nolint:gosec // target — только внутренний путь: его проверил safeRedirect
|
||
|
|
http.Redirect(w, r, target+separator+"oauth=ok", http.StatusFound)
|
||
|
|
}
|
||
|
|
|
||
|
|
// oauthErrorCode переводит доменную ошибку в код для `?oauth_error=`.
|
||
|
|
func oauthErrorCode(err error) string {
|
||
|
|
for _, candidate := range []error{
|
||
|
|
auth.ErrOAuthNotConfigured,
|
||
|
|
auth.ErrOAuthUnknownProvider,
|
||
|
|
auth.ErrOAuthState,
|
||
|
|
auth.ErrOAuthEmailUnverified,
|
||
|
|
auth.ErrOAuthEmailMissing,
|
||
|
|
auth.ErrRegistrationOff,
|
||
|
|
auth.ErrUserBanned,
|
||
|
|
auth.ErrOAuthExchange,
|
||
|
|
} {
|
||
|
|
if errors.Is(err, candidate) {
|
||
|
|
return strings.TrimPrefix(candidate.Error(), "oauth.")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return "failed"
|
||
|
|
}
|