Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 75 additions & 0 deletions apps/sim/app/api/auth/sso/providers/route.test.ts
Original file line number Diff line number Diff line change
@@ -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)
})
})
146 changes: 65 additions & 81 deletions apps/sim/app/api/auth/sso/providers/route.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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'

Expand All @@ -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) {
Comment thread
waleedlatif1 marked this conversation as resolved.
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,
Comment thread
waleedlatif1 marked this conversation as resolved.
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 })
Expand Down
14 changes: 4 additions & 10 deletions apps/sim/app/api/auth/sso/resolve/route.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -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' })
)
Expand All @@ -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' }))
Expand Down
17 changes: 7 additions & 10 deletions apps/sim/app/api/auth/sso/resolve/route.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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(
Expand All @@ -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 })
})
4 changes: 2 additions & 2 deletions apps/sim/ee/sso/components/sso-form.test.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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')

Expand Down
31 changes: 14 additions & 17 deletions apps/sim/lib/api/contracts/auth.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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({
Expand Down Expand Up @@ -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() }),
},
})

Expand Down
Loading