diff --git a/web/src/app/router.tsx b/web/src/app/router.tsx index d20ea269..216c4377 100644 --- a/web/src/app/router.tsx +++ b/web/src/app/router.tsx @@ -4,7 +4,7 @@ import { Layout } from './layout' import { getCurrentUser } from '@/api/client' import { RoleGuard } from '@/shared/components/role-guard' import { RouteError } from '@/shared/components/route-error' -import { createRequireAuth } from '@/shared/lib/auth-route' +import { createRequireAuth, isSafeAuthReturnTo } from '@/shared/lib/auth-route' import { clearDynamicImportReloadGuard, recoverFromDynamicImportError } from '@/shared/lib/dynamic-import-recovery' import { normalizeSearchQuery } from '@/shared/lib/search-query' @@ -245,7 +245,7 @@ const loginRoute = createRoute({ getParentRoute: () => rootRoute, path: 'login', validateSearch: (search: Record): { returnTo?: string; reason?: string } => ({ - returnTo: typeof search.returnTo === 'string' && search.returnTo ? search.returnTo : undefined, + returnTo: isSafeAuthReturnTo(search.returnTo) ? search.returnTo : undefined, reason: typeof search.reason === 'string' ? search.reason : undefined, }), component: LoginPage, @@ -255,7 +255,7 @@ const registerRoute = createRoute({ getParentRoute: () => rootRoute, path: 'register', validateSearch: (search: Record) => ({ - returnTo: typeof search.returnTo === 'string' ? search.returnTo : '', + returnTo: isSafeAuthReturnTo(search.returnTo) ? search.returnTo : '', }), component: RegisterPage, }) diff --git a/web/src/pages/login.test.tsx b/web/src/pages/login.test.tsx index 7d26202c..add82cac 100644 --- a/web/src/pages/login.test.tsx +++ b/web/src/pages/login.test.tsx @@ -1,5 +1,5 @@ /** @vitest-environment jsdom */ -import { cleanup, fireEvent, render, screen } from '@testing-library/react' +import { cleanup, fireEvent, render, screen, waitFor } from '@testing-library/react' import { afterEach } from 'vitest' import { describe, expect, it, vi } from 'vitest' @@ -7,12 +7,15 @@ const authMethodsFixture = vi.hoisted(() => ({ methods: [{ id: 'organization', methodType: 'ENTERPRISE_DISCOVERY' }], bootstrapEnabled: false, isError: false, + returnTo: '', + navigate: vi.fn(), + mutateAsync: vi.fn(), })) vi.mock('@tanstack/react-router', () => ({ Link: ({ children }: { children: unknown }) => children, - useNavigate: () => vi.fn(), - useSearch: () => ({ returnTo: '' }), + useNavigate: () => authMethodsFixture.navigate, + useSearch: () => ({ returnTo: authMethodsFixture.returnTo }), })) vi.mock('react-i18next', async () => { @@ -62,7 +65,7 @@ vi.mock('@/features/auth/use-auth-methods', () => ({ vi.mock('@/features/auth/use-password-login', () => ({ usePasswordLogin: () => ({ - mutateAsync: vi.fn(), + mutateAsync: authMethodsFixture.mutateAsync, isPending: false, error: null, }), @@ -85,6 +88,9 @@ describe('LoginPage', () => { authMethodsFixture.methods = [{ id: 'organization', methodType: 'ENTERPRISE_DISCOVERY' }] authMethodsFixture.bootstrapEnabled = false authMethodsFixture.isError = false + authMethodsFixture.returnTo = '' + authMethodsFixture.navigate.mockClear() + authMethodsFixture.mutateAsync.mockClear() }) it('exports a named component function', () => { @@ -102,6 +108,23 @@ describe('LoginPage', () => { expect(html).toContain('login.register') }) + it('returns to the home page after direct login without an explicit destination', async () => { + render() + fireEvent.change(screen.getByLabelText('login.username'), { target: { value: 'user1' } }) + fireEvent.change(screen.getByLabelText('login.password'), { target: { value: 'password' } }) + fireEvent.click(screen.getByRole('button', { name: 'login.submit' })) + await waitFor(() => expect(authMethodsFixture.navigate).toHaveBeenCalledWith({ to: '/' })) + }) + + it('returns to the original local page after direct login', async () => { + authMethodsFixture.returnTo = '/skills?tab=mine' + render() + fireEvent.change(screen.getByLabelText('login.username'), { target: { value: 'user1' } }) + fireEvent.change(screen.getByLabelText('login.password'), { target: { value: 'password' } }) + fireEvent.click(screen.getByRole('button', { name: 'login.submit' })) + await waitFor(() => expect(authMethodsFixture.navigate).toHaveBeenCalledWith({ to: '/skills?tab=mine' })) + }) + it('replaces the password form with organization discovery and can switch back', () => { render() const personal = screen.getByRole('button', { name: 'login.tabPersonal' }) diff --git a/web/src/pages/login.tsx b/web/src/pages/login.tsx index 9e9307a6..658014db 100644 --- a/web/src/pages/login.tsx +++ b/web/src/pages/login.tsx @@ -11,6 +11,7 @@ import { useAuthMethods } from '@/features/auth/use-auth-methods' import { usePasswordLogin } from '@/features/auth/use-password-login' import { Button } from '@/shared/ui/button' import { Input } from '@/shared/ui/input' +import { resolveAuthReturnTo } from '@/shared/lib/auth-route' /** * Authentication entry page. @@ -31,9 +32,8 @@ export function LoginPage() { const [loginMode, setLoginMode] = useState<'personal' | 'organization'>('personal') const [fieldErrors, setFieldErrors] = useState<{ username?: string, password?: string }>({}) const isChinese = i18n.resolvedLanguage?.split('-')[0] === 'zh' - const { data: authMethods, isLoading: authMethodsLoading } = useAuthMethods(search.returnTo) - - const returnTo = search.returnTo && search.returnTo.startsWith('/') ? search.returnTo : '/dashboard' + const returnTo = resolveAuthReturnTo(search.returnTo) + const { data: authMethods, isLoading: authMethodsLoading } = useAuthMethods(returnTo) const disabledMessage = search.reason === 'accountDisabled' ? t('apiError.auth.accountDisabled') : null const directMethod = directAuthConfig.provider ? authMethods?.find((method) => diff --git a/web/src/pages/register.tsx b/web/src/pages/register.tsx index b580926d..18fc234e 100644 --- a/web/src/pages/register.tsx +++ b/web/src/pages/register.tsx @@ -7,6 +7,7 @@ import { LoginButton } from '@/features/auth/login-button' import { useLocalRegister } from '@/features/auth/use-local-auth' import { Button } from '@/shared/ui/button' import { Input } from '@/shared/ui/input' +import { resolveAuthReturnTo } from '@/shared/lib/auth-route' const USERNAME_PATTERN = /^[A-Za-z0-9_]{3,64}$/ const EMAIL_PATTERN = /^[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Za-z]{2,}$/ @@ -60,7 +61,7 @@ export function RegisterPage() { const [fieldErrors, setFieldErrors] = useState({}) const [formError, setFormError] = useState(null) - const returnTo = search.returnTo && search.returnTo.startsWith('/') ? search.returnTo : '/dashboard' + const returnTo = resolveAuthReturnTo(search.returnTo) function validateUsername(value: string) { const trimmed = value.trim() diff --git a/web/src/shared/lib/auth-route.test.ts b/web/src/shared/lib/auth-route.test.ts index c15eb46a..695e6dff 100644 --- a/web/src/shared/lib/auth-route.test.ts +++ b/web/src/shared/lib/auth-route.test.ts @@ -1,8 +1,16 @@ import { describe, expect, it, vi } from 'vitest' import { isRedirect } from '@tanstack/react-router' -import { buildReturnTo, createRequireAuth } from './auth-route' +import { buildReturnTo, createRequireAuth, resolveAuthReturnTo } from './auth-route' describe('auth-route', () => { + it('returns to the original local page or the home page, never an external URL', () => { + expect(resolveAuthReturnTo('/skills?tab=mine#latest')).toBe('/skills?tab=mine#latest') + expect(resolveAuthReturnTo(undefined)).toBe('/') + expect(resolveAuthReturnTo('//example.com')).toBe('/') + expect(resolveAuthReturnTo('/\\example.com')).toBe('/') + expect(resolveAuthReturnTo('https://example.com')).toBe('/') + }) + it('buildReturnTo preserves pathname search and hash', () => { expect(buildReturnTo({ pathname: '/space/global/caldav-calendar', diff --git a/web/src/shared/lib/auth-route.ts b/web/src/shared/lib/auth-route.ts index df047ec8..492be432 100644 --- a/web/src/shared/lib/auth-route.ts +++ b/web/src/shared/lib/auth-route.ts @@ -10,6 +10,17 @@ export function buildReturnTo(location: RouteLocationLike) { return `${location.pathname}${location.searchStr ?? ''}${location.hash ?? ''}` } +export function isSafeAuthReturnTo(value: unknown): value is string { + return typeof value === 'string' + && value.startsWith('/') + && !value.startsWith('//') + && !value.includes('\\') +} + +export function resolveAuthReturnTo(value: unknown) { + return isSafeAuthReturnTo(value) ? value : '/' +} + export function createRequireAuth(getCurrentUser: () => Promise) { return async function requireAuth({ location }: { location: RouteLocationLike }) { const user = await getCurrentUser()