diff --git a/backend/biz/user/handler/v1/auth.go b/backend/biz/user/handler/v1/auth.go index fded268e..14157924 100644 --- a/backend/biz/user/handler/v1/auth.go +++ b/backend/biz/user/handler/v1/auth.go @@ -4,12 +4,15 @@ import ( "fmt" "log/slog" "net/http" + "net/url" + "time" "github.com/GoYoko/web" "github.com/google/uuid" "github.com/redis/go-redis/v9" "github.com/samber/do" + "github.com/chaitin/MonkeyCode/backend/biz/user/provider" "github.com/chaitin/MonkeyCode/backend/config" "github.com/chaitin/MonkeyCode/backend/consts" "github.com/chaitin/MonkeyCode/backend/domain" @@ -31,6 +34,8 @@ type AuthHandler struct { captcha *captcha.Captcha } +const oidcLoginRedirectKeyPrefix = "oidc_login_redirect:" + // NewAuthHandler 创建认证处理器 (samber/do 风格) func NewAuthHandler(i *do.Injector) (*AuthHandler, error) { w := do.MustInvoke[*web.Web](i) @@ -98,10 +103,26 @@ func (h *AuthHandler) OIDCLogin(c *web.Context, req domain.TeamOIDCLoginReq) err if h.oidcUsecase == nil { return errcode.ErrOIDCDisabled } + redirectURL, err := provider.CleanRedirectURL(req.RedirectURL) + if err != nil { + return errcode.ErrOAuthLoginRedirectInvalid + } authURL, err := h.oidcUsecase.StartLogin(c.Request().Context(), req.TeamID) if err != nil { return err } + parsedAuthURL, err := url.Parse(authURL) + if err != nil || parsedAuthURL.Query().Get("state") == "" { + return errcode.ErrInternalServer + } + if err := h.redis.Set( + c.Request().Context(), + oidcLoginRedirectKeyPrefix+parsedAuthURL.Query().Get("state"), + redirectURL, + 10*time.Minute, + ).Err(); err != nil { + return errcode.ErrInternalServer + } return c.Redirect(http.StatusFound, authURL) } @@ -119,16 +140,26 @@ func (h *AuthHandler) OIDCCallback(c *web.Context, req domain.TeamOIDCCallbackRe if h.oidcUsecase == nil { return errcode.ErrOIDCDisabled } - user, err := h.oidcUsecase.HandleCallback(c.Request().Context(), &req) + ctx := c.Request().Context() + redirectURL, err := h.redis.Get(ctx, oidcLoginRedirectKeyPrefix+req.State).Result() + if err == redis.Nil { + redirectURL = "/console/" + } else if err != nil { + return errcode.ErrInternalServer + } + user, err := h.oidcUsecase.HandleCallback(ctx, &req) if err != nil { return err } _, err = h.authMiddleware.Session.Save(c, consts.MonkeyCodeAISession, user.ID, user) if err != nil { - h.logger.ErrorContext(c.Request().Context(), "save oidc session failed", "error", err) + h.logger.ErrorContext(ctx, "save oidc session failed", "error", err) return errcode.ErrInternalServer } - return c.Redirect(http.StatusFound, "/console/") + if err := h.redis.Del(ctx, oidcLoginRedirectKeyPrefix+req.State).Err(); err != nil { + h.logger.WarnContext(ctx, "delete oidc login redirect failed", "error", err) + } + return c.Redirect(http.StatusFound, redirectURL) } // OIDCPublicConfig 获取团队公开 OIDC 登录配置 diff --git a/backend/domain/team_oidc.go b/backend/domain/team_oidc.go index befe33fd..de982737 100644 --- a/backend/domain/team_oidc.go +++ b/backend/domain/team_oidc.go @@ -81,7 +81,8 @@ type TeamOIDCPublicConfigResp struct { } type TeamOIDCLoginReq struct { - TeamID uuid.UUID `query:"team_id" validate:"required"` + TeamID uuid.UUID `query:"team_id" validate:"required"` + RedirectURL string `query:"redirect_url"` } type TeamOIDCCallbackReq struct { diff --git a/frontend/src/pages/login.tsx b/frontend/src/pages/login.tsx index 37af262d..49515d63 100644 --- a/frontend/src/pages/login.tsx +++ b/frontend/src/pages/login.tsx @@ -21,7 +21,7 @@ import { Spinner } from "@/components/ui/spinner" import React from "react" import { toast } from "sonner" import { apiRequest } from "@/utils/requestUtils" -import { Link, useNavigate } from "react-router-dom" +import { Link, useNavigate, useSearchParams } from "react-router-dom" import { captchaChallenge } from "@/utils/common" import { ArrowLeft, Eye, EyeOff } from "lucide-react" import { IconBrandGithub, IconBrandGoogle } from "@tabler/icons-react" @@ -33,8 +33,17 @@ import { useAppRuntime } from "@/components/app-runtime-provider" const USER_STORAGE_KEY = 'login_user' const MANAGER_STORAGE_KEY = 'login_manager' +const DEFAULT_USER_REDIRECT = '/console/tasks' type OAuthProvider = "github" | "google" +// 仅接受站内绝对路径,避免 redir 被用于跳转到外部站点。 +function sanitizeRedirect(raw: string | null): string { + if (!raw) return '' + if (/[\r\n\t\\]/.test(raw)) return '' + if (!raw.startsWith('/') || raw.startsWith('//')) return '' + return raw +} + export default function LoginPage({ className, ...props @@ -51,14 +60,19 @@ export default function LoginPage({ const [oauthLoggingProvider, setOauthLoggingProvider] = React.useState(null) const [defaultOIDCConfig, setDefaultOIDCConfig] = React.useState(null) const navigate = useNavigate() + const [searchParams] = useSearchParams() + const userRedirect = sanitizeRedirect(searchParams.get('redir')) || DEFAULT_USER_REDIRECT const { t } = useTranslation() const { captchaEnabled, reloadAuth, serverConfig } = useAppRuntime() const serverRegion = serverConfig?.region as string | undefined const isCnRegion = serverRegion === "cn" const isGlobalRegion = serverRegion === "global" const inviterId = typeof window !== 'undefined' ? (localStorage.getItem('ic') || '') : '' - const userLoginHref = `/api/v1/users/login?redirect=&inviter_id=${inviterId}` + const userLoginHref = `/api/v1/users/login?redirect=${encodeURIComponent(userRedirect)}&inviter_id=${inviterId}` const defaultOIDCLoginURL = defaultOIDCConfig?.enabled ? defaultOIDCConfig.login_url : '' + const defaultOIDCLoginHref = defaultOIDCLoginURL + ? `${defaultOIDCLoginURL}${defaultOIDCLoginURL.includes('?') ? '&' : '?'}redirect_url=${encodeURIComponent(userRedirect)}` + : '' const ensureTermsAccepted = React.useCallback(() => { if (agreedToTerms) return true @@ -126,7 +140,7 @@ export default function LoginPage({ if (resp.code === 0) { localStorage.setItem(USER_STORAGE_KEY, JSON.stringify({ email: userEmail.trim(), password: userPassword.trim() })) await reloadAuth() - navigate('/console/tasks') + window.location.assign(userRedirect) } else { toast.error(t("login.toast.loginFailed")) } @@ -145,7 +159,7 @@ export default function LoginPage({ try { const response = await new Api().api.v1UsersOauthLoginDetail(provider, { - redirect_url: "/console/tasks", + redirect_url: userRedirect, }) const authUrl = response.data?.data?.auth_url @@ -259,7 +273,7 @@ export default function LoginPage({ {IS_OFFLINE_EDITION && defaultOIDCLoginURL && (