diff --git a/apps/sim/app/api/auth/sso/providers/route.test.ts b/apps/sim/app/api/auth/sso/providers/route.test.ts new file mode 100644 index 00000000000..47c7cfdcae1 --- /dev/null +++ b/apps/sim/app/api/auth/sso/providers/route.test.ts @@ -0,0 +1,75 @@ +/** + * @vitest-environment node + */ +import { + createMockRequest, + dbChainMock, + dbChainMockFns, + queueTableRows, + resetDbChainMock, + schemaMock, +} from '@sim/testing' +import { beforeEach, describe, expect, it, vi } from 'vitest' + +const { mockGetSession } = vi.hoisted(() => ({ mockGetSession: vi.fn() })) + +vi.mock('@sim/db', () => ({ ...dbChainMock, ...schemaMock })) +vi.mock('@/lib/auth', () => ({ getSession: mockGetSession })) + +import { GET } from '@/app/api/auth/sso/providers/route' + +const providerRow = { + id: 'row-1', + providerId: 'acme-okta', + domain: 'acme.com', + issuer: 'https://acme.okta.test', + oidcConfig: JSON.stringify({ clientId: 'client', clientSecret: 'a-long-client-secret-wxyz' }), + samlConfig: null, + userId: 'user-1', + organizationId: 'org-1', + jitProvisioningEnabled: true, + domainVerified: true, + domainKey: 'acme.com', + isNamedPrimary: false, +} + +describe('GET /api/auth/sso/providers', () => { + beforeEach(() => { + vi.clearAllMocks() + resetDbChainMock() + mockGetSession.mockResolvedValue({ user: { id: 'user-1' } }) + }) + + it('refuses a caller without a session before reading any provider', async () => { + mockGetSession.mockResolvedValue(null) + const res = await GET(createMockRequest('GET')) + expect(res.status).toBe(401) + expect(dbChainMockFns.select).not.toHaveBeenCalled() + }) + + it('lists only the providers the caller registered when no organization is named', async () => { + queueTableRows(schemaMock.ssoProvider, [providerRow]) + const res = await GET(createMockRequest('GET')) + expect(res.status).toBe(200) + const { providers } = await res.json() + expect(providers).toHaveLength(1) + expect(providers[0]).toMatchObject({ providerId: 'acme-okta', providerType: 'oidc' }) + expect(JSON.parse(providers[0].oidcConfig)).toMatchObject({ clientSecretHint: 'wxyz' }) + expect(providers[0].oidcConfig).not.toContain('a-long-client-secret') + const condition = JSON.stringify(dbChainMockFns.where.mock.calls[0][0]) + expect(condition).toContain('user-1') + }) + + it('refuses an organization the caller does not administer', async () => { + queueTableRows(schemaMock.member, [{ organizationId: 'org-1', role: 'member' }]) + const res = await GET( + createMockRequest( + 'GET', + undefined, + {}, + 'http://localhost/api/auth/sso/providers?organizationId=org-1' + ) + ) + expect(res.status).toBe(403) + }) +}) diff --git a/apps/sim/app/api/auth/sso/providers/route.ts b/apps/sim/app/api/auth/sso/providers/route.ts index e2f2767dc5e..9a447da9df0 100644 --- a/apps/sim/app/api/auth/sso/providers/route.ts +++ b/apps/sim/app/api/auth/sso/providers/route.ts @@ -11,7 +11,6 @@ import { listSsoProvidersContract } from '@/lib/api/contracts/auth' import { parseRequest } from '@/lib/api/server' import { getSession } from '@/lib/auth' import { markSignInProviders } from '@/lib/auth/sso/primary-provider' -import { enforceIpRateLimit } from '@/lib/core/rate-limiter' import { REDACTED_MARKER } from '@/lib/core/security/redaction' import { withRouteHandler } from '@/lib/core/utils/with-route-handler' @@ -32,102 +31,87 @@ function buildClientSecretHint(clientSecret: unknown): string | null { return clientSecret.slice(-4) } +/** + * Lists the identity providers the caller administers: an organization's when an + * owner or admin names it, otherwise the ones the caller registered. + * + * Signed-in only. Sign-in resolves one address at a time through + * `/api/auth/sso/resolve`; nothing needs every configured domain, and listing + * them would publish which organizations use SSO. + */ export const GET = withRouteHandler(async (request: NextRequest) => { try { const session = await getSession() if (!session?.user?.id) { - const rateLimited = await enforceIpRateLimit('sso-providers', request, { - maxTokens: 20, - refillRate: 20, - refillIntervalMs: 60_000, - }) - if (rateLimited) return rateLimited + return NextResponse.json({ error: 'Unauthorized' }, { status: 401 }) } const parsed = await parseRequest(listSsoProvidersContract, request, {}) if (!parsed.success) return parsed.response const { organizationId } = parsed.data.query + const userId = session.user.id - let providers - if (session?.user?.id) { - const userId = session.user.id - - let verifiedOrganizationId: string | null = null - if (organizationId) { - const [membership] = await db - .select({ organizationId: member.organizationId, role: member.role }) - .from(member) - .where(and(eq(member.userId, userId), eq(member.organizationId, organizationId))) - .limit(1) - if (!membership) { - return NextResponse.json({ error: 'Forbidden' }, { status: 403 }) - } - if (membership.role !== 'owner' && membership.role !== 'admin') { - return NextResponse.json({ error: 'Forbidden' }, { status: 403 }) - } - verifiedOrganizationId = membership.organizationId + let verifiedOrganizationId: string | null = null + if (organizationId) { + const [membership] = await db + .select({ organizationId: member.organizationId, role: member.role }) + .from(member) + .where(and(eq(member.userId, userId), eq(member.organizationId, organizationId))) + .limit(1) + if (!membership) { + return NextResponse.json({ error: 'Forbidden' }, { status: 403 }) } + if (membership.role !== 'owner' && membership.role !== 'admin') { + return NextResponse.json({ error: 'Forbidden' }, { status: 403 }) + } + verifiedOrganizationId = membership.organizationId + } - const whereClause = verifiedOrganizationId - ? eq(ssoProvider.organizationId, verifiedOrganizationId) - : eq(ssoProvider.userId, userId) - - const results = await db - .select({ - id: ssoProvider.id, - providerId: ssoProvider.providerId, - domain: ssoProvider.domain, - issuer: ssoProvider.issuer, - oidcConfig: ssoProvider.oidcConfig, - samlConfig: ssoProvider.samlConfig, - userId: ssoProvider.userId, - organizationId: ssoProvider.organizationId, - jitProvisioningEnabled: ssoProvider.jitProvisioningEnabled, - domainVerified: ssoProvider.domainVerified, - domainKey: ssoProviderDomainKey, - isNamedPrimary, - }) - .from(ssoProvider) - .leftJoin(ssoDomain, verifiedDomainOfProvider) - .where(whereClause) - .orderBy(asc(ssoProvider.providerId)) + const whereClause = verifiedOrganizationId + ? eq(ssoProvider.organizationId, verifiedOrganizationId) + : eq(ssoProvider.userId, userId) - providers = markSignInProviders(results).map((provider) => { - let oidcConfig = provider.oidcConfig - if (oidcConfig) { - try { - const parsed = JSON.parse(oidcConfig) - const hint = buildClientSecretHint(parsed.clientSecret) - parsed.clientSecret = REDACTED_MARKER - if (hint) parsed.clientSecretHint = hint - oidcConfig = JSON.stringify(parsed) - } catch { - oidcConfig = null - } - } - return { - ...provider, - oidcConfig, - providerType: (provider.samlConfig ? 'saml' : 'oidc') as 'oidc' | 'saml', - } + const results = await db + .select({ + id: ssoProvider.id, + providerId: ssoProvider.providerId, + domain: ssoProvider.domain, + issuer: ssoProvider.issuer, + oidcConfig: ssoProvider.oidcConfig, + samlConfig: ssoProvider.samlConfig, + userId: ssoProvider.userId, + organizationId: ssoProvider.organizationId, + jitProvisioningEnabled: ssoProvider.jitProvisioningEnabled, + domainVerified: ssoProvider.domainVerified, + domainKey: ssoProviderDomainKey, + isNamedPrimary, }) - } else { - const results = await db - .select({ - domain: ssoProvider.domain, - }) - .from(ssoProvider) - - providers = results.map((provider) => ({ - domain: provider.domain, - })) - } + .from(ssoProvider) + .leftJoin(ssoDomain, verifiedDomainOfProvider) + .where(whereClause) + .orderBy(asc(ssoProvider.providerId)) - logger.info('Fetched SSO providers', { - userId: session?.user?.id, - authenticated: !!session?.user?.id, - providerCount: providers.length, + const providers = markSignInProviders(results).map((provider) => { + let oidcConfig = provider.oidcConfig + if (oidcConfig) { + try { + const parsed = JSON.parse(oidcConfig) + const hint = buildClientSecretHint(parsed.clientSecret) + parsed.clientSecret = REDACTED_MARKER + if (hint) parsed.clientSecretHint = hint + oidcConfig = JSON.stringify(parsed) + } catch { + oidcConfig = null + } + } + return { + ...provider, + oidcConfig, + providerType: (provider.samlConfig ? 'saml' : 'oidc') as 'oidc' | 'saml', + } }) + logger.info('Fetched SSO providers', { userId, providerCount: providers.length }) + return NextResponse.json({ providers }) } catch (error) { logger.error('Failed to fetch SSO providers', { error }) diff --git a/apps/sim/app/api/auth/sso/resolve/route.test.ts b/apps/sim/app/api/auth/sso/resolve/route.test.ts index 90567f03bbd..0c1cf725ccd 100644 --- a/apps/sim/app/api/auth/sso/resolve/route.test.ts +++ b/apps/sim/app/api/auth/sso/resolve/route.test.ts @@ -26,17 +26,17 @@ describe('POST /api/auth/sso/resolve', () => { }) it('names the provider that serves the address domain', async () => { - queueTableRows(schemaMock.ssoProvider, [{ providerId: 'acme-okta', samlConfig: null }]) + queueTableRows(schemaMock.ssoProvider, [{ providerId: 'acme-okta' }]) const res = await POST(createMockRequest('POST', { email: 'Ada@Acme.com' })) expect(res.status).toBe(200) - await expect(res.json()).resolves.toEqual({ providerId: 'acme-okta', providerType: 'oidc' }) + await expect(res.json()).resolves.toEqual({ providerId: 'acme-okta' }) const [condition] = dbChainMockFns.where.mock.calls[0] expect(JSON.stringify(condition)).toContain('acme.com') expect(JSON.stringify(condition)).toContain('domainVerified') }) it('prefers the provider the verified domain names, then provider id', async () => { - queueTableRows(schemaMock.ssoProvider, [{ providerId: 'acme-okta', samlConfig: null }]) + queueTableRows(schemaMock.ssoProvider, [{ providerId: 'acme-okta' }]) await POST(createMockRequest('POST', { email: 'ada@acme.com' })) expect(dbChainMockFns.leftJoin).toHaveBeenCalledWith(schemaMock.ssoDomain, expect.anything()) const [named, byId] = dbChainMockFns.orderBy.mock.calls[0] @@ -46,7 +46,7 @@ describe('POST /api/auth/sso/resolve', () => { }) it('honors a test link only for a provider that serves the address domain', async () => { - queueTableRows(schemaMock.ssoProvider, [{ providerId: 'acme-entra', samlConfig: null }]) + queueTableRows(schemaMock.ssoProvider, [{ providerId: 'acme-entra' }]) const res = await POST( createMockRequest('POST', { email: 'ada@acme.com', providerId: 'acme-entra' }) ) @@ -65,12 +65,6 @@ describe('POST /api/auth/sso/resolve', () => { expect(res.status).toBe(404) }) - it('reports SAML providers as such', async () => { - queueTableRows(schemaMock.ssoProvider, [{ providerId: 'acme-adfs', samlConfig: '{}' }]) - const res = await POST(createMockRequest('POST', { email: 'ada@acme.com' })) - await expect(res.json()).resolves.toMatchObject({ providerType: 'saml' }) - }) - it('answers 404 when no provider serves the domain', async () => { queueTableRows(schemaMock.ssoProvider, []) const res = await POST(createMockRequest('POST', { email: 'ada@nowhere.test' })) diff --git a/apps/sim/app/api/auth/sso/resolve/route.ts b/apps/sim/app/api/auth/sso/resolve/route.ts index bd038bee469..8f5b0480043 100644 --- a/apps/sim/app/api/auth/sso/resolve/route.ts +++ b/apps/sim/app/api/auth/sso/resolve/route.ts @@ -16,11 +16,11 @@ import { withRouteHandler } from '@/lib/core/utils/with-route-handler' * Names the identity provider that signs in an email address. * * Unauthenticated by nature, like the sign-in page that calls it, and admitted - * per address. It discloses nothing the public provider list does not: which - * domains have SSO, and the provider id that already appears in the callback - * URL. Only a provider whose domain is verified is named: an unverified claim - * has no authority over the address, and sending someone to its IdP would fail - * at the callback anyway. + * per address. It answers for the one domain asked about, and names only the + * provider id that the sign-in redirect and callback URL expose anyway; there is + * deliberately no way to list every domain with SSO. Only a provider whose domain + * is verified is named: an unverified claim has no authority over the address, + * and sending someone to its IdP would fail at the callback anyway. * * The provider the domain names as primary wins, then the first verified by id, * which is also the only one when a domain has a single provider. A test sign-in @@ -46,7 +46,7 @@ export const POST = withRouteHandler(async (request: NextRequest) => { const requestedProviderId = parsed.data.body.providerId const [provider] = await db - .select({ providerId: ssoProvider.providerId, samlConfig: ssoProvider.samlConfig }) + .select({ providerId: ssoProvider.providerId }) .from(ssoProvider) .leftJoin(ssoDomain, verifiedDomainOfProvider) .where( @@ -68,8 +68,5 @@ export const POST = withRouteHandler(async (request: NextRequest) => { ) } - return NextResponse.json({ - providerId: provider.providerId, - providerType: provider.samlConfig ? 'saml' : 'oidc', - }) + return NextResponse.json({ providerId: provider.providerId }) }) diff --git a/apps/sim/ee/sso/components/sso-form.test.tsx b/apps/sim/ee/sso/components/sso-form.test.tsx index 08cd367b44d..37f92d89d94 100644 --- a/apps/sim/ee/sso/components/sso-form.test.tsx +++ b/apps/sim/ee/sso/components/sso-form.test.tsx @@ -144,7 +144,7 @@ describe('SSOForm sign-in errors', () => { mockSsoSignIn.mockReset() mockUseSearchParams.mockReset() mockRequestJson.mockReset() - mockRequestJson.mockResolvedValue({ providerId: 'example-okta', providerType: 'oidc' }) + mockRequestJson.mockResolvedValue({ providerId: 'example-okta' }) ;(globalThis as { IS_REACT_ACT_ENVIRONMENT?: boolean }).IS_REACT_ACT_ENVIRONMENT = true container = document.createElement('div') document.body.appendChild(container) @@ -188,7 +188,7 @@ describe('SSOForm sign-in errors', () => { }) it('signs in through the provider a test link names, and returns to that link on failure', async () => { - mockRequestJson.mockResolvedValue({ providerId: 'example-okta', providerType: 'oidc' }) + mockRequestJson.mockResolvedValue({ providerId: 'example-okta' }) mockSsoSignIn.mockResolvedValue({ data: { url: 'https://idp.example.com' }, error: null }) renderInteractive('email=user%40example.com&provider=example-okta') diff --git a/apps/sim/lib/api/contracts/auth.ts b/apps/sim/lib/api/contracts/auth.ts index f8cd13f45f8..cdd96d38547 100644 --- a/apps/sim/lib/api/contracts/auth.ts +++ b/apps/sim/lib/api/contracts/auth.ts @@ -95,22 +95,22 @@ export const ssoRegistrationContract = defineRouteContract({ }) const ssoProviderListEntrySchema = z.object({ - id: z.string().optional(), - providerId: z.string().optional(), - domain: z.string().nullable(), - issuer: z.string().nullable().optional(), - oidcConfig: z.string().nullable().optional(), - samlConfig: z.string().nullable().optional(), - userId: z.string().nullable().optional(), - organizationId: z.string().nullable().optional(), - jitProvisioningEnabled: z.boolean().optional(), + id: z.string(), + providerId: z.string(), + domain: z.string(), + issuer: z.string(), + oidcConfig: z.string().nullable(), + samlConfig: z.string().nullable(), + userId: z.string(), + organizationId: z.string().nullable(), + jitProvisioningEnabled: z.boolean(), /** The domain as sign-in compares it: trimmed, lower-cased, a leading `*.` dropped. Providers sharing it share a primary. */ - domainKey: z.string().optional(), + domainKey: z.string(), /** Whether this provider's domain is verified, so it can sign people in. */ - domainVerified: z.boolean().optional(), + domainVerified: z.boolean(), /** Whether sign-in for this provider's domain goes through it. */ - isPrimary: z.boolean().optional(), - providerType: z.enum(['oidc', 'saml']).optional(), + isPrimary: z.boolean(), + providerType: z.enum(['oidc', 'saml']), }) export const listSsoProvidersContract = defineRouteContract({ @@ -167,10 +167,7 @@ export const resolveSsoProviderContract = defineRouteContract({ }), response: { mode: 'json', - schema: z.object({ - providerId: z.string(), - providerType: z.enum(['oidc', 'saml']), - }), + schema: z.object({ providerId: z.string() }), }, })