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
129 changes: 58 additions & 71 deletions apps/sim/app/api/auth/sso/register/route.test.ts
Original file line number Diff line number Diff line change
@@ -1,67 +1,46 @@
/**
* @vitest-environment node
*/
import { createEnvMock, createMockRequest } from '@sim/testing'
import { beforeEach, describe, expect, it, vi } from 'vitest'
import {
createEnvMock,
createMockRequest,
dbChainMock,
dbChainMockFns,
queueTableRows,
resetDbChainMock,
schemaMock,
} from '@sim/testing'
import { afterAll, beforeEach, describe, expect, it, vi } from 'vitest'

const {
mockGetSession,
mockRegisterSSOProvider,
mockHasSSOAccess,
mockValidateUrlWithDNS,
mockSecureFetchWithPinnedIP,
dbState,
memberTable,
ssoProviderTable,
} = vi.hoisted(() => ({
mockGetSession: vi.fn(),
mockRegisterSSOProvider: vi.fn(),
mockHasSSOAccess: vi.fn(),
mockValidateUrlWithDNS: vi.fn(),
mockSecureFetchWithPinnedIP: vi.fn(),
dbState: { members: [] as any[], providers: [] as any[] },
memberTable: {
userId: 'member.userId',
organizationId: 'member.organizationId',
role: 'member.role',
},
ssoProviderTable: {
id: 'sso.id',
providerId: 'sso.providerId',
domain: 'sso.domain',
issuer: 'sso.issuer',
userId: 'sso.userId',
organizationId: 'sso.organizationId',
oidcConfig: 'sso.oidcConfig',
samlConfig: 'sso.samlConfig',
},
}))

function makeBuilder(rows: any[]): any {
const thenable: any = Promise.resolve(rows)
thenable.where = (condition: any) => {
const values = condition?.values
if (Array.isArray(values) && values.length > 0) {
const target = String(values[values.length - 1]).toLowerCase()
return makeBuilder(rows.filter((r) => String(r.domain ?? '').toLowerCase() === target))
}
return makeBuilder(rows)
}
thenable.limit = () => Promise.resolve(rows)
thenable.orderBy = () => Promise.resolve(rows)
return thenable
vi.mock('@sim/db', () => ({ ...dbChainMock, ...schemaMock }))

/** Queues the caller's org membership row(s) for the admin/owner check. */
function queueMembers(rows: Array<Record<string, unknown>>) {
queueTableRows(schemaMock.member, rows)
}

vi.mock('@sim/db', () => ({
db: {
select: () => ({
from: (table: unknown) =>
makeBuilder(table === memberTable ? dbState.members : dbState.providers),
}),
},
member: memberTable,
ssoProvider: ssoProviderTable,
}))
/**
* Queues existing SSO provider rows for BOTH domain-conflict lookups (the
* pre-registration check and the post-registration re-check).
*/
function queueProviders(rows: Array<Record<string, unknown>>) {
queueTableRows(schemaMock.ssoProvider, rows)
queueTableRows(schemaMock.ssoProvider, rows)
}

vi.mock('@/lib/auth', () => ({
getSession: mockGetSession,
Expand Down Expand Up @@ -109,15 +88,18 @@ function request(body: Record<string, unknown>) {
describe('POST /api/auth/sso/register', () => {
beforeEach(() => {
vi.clearAllMocks()
dbState.members = []
dbState.providers = []
resetDbChainMock()
mockGetSession.mockResolvedValue({ user: { id: 'u1' } })
mockHasSSOAccess.mockResolvedValue(true)
mockValidateUrlWithDNS.mockResolvedValue({ isValid: true, resolvedIP: '1.2.3.4' })
mockSecureFetchWithPinnedIP.mockRejectedValue(new Error('discovery not mocked for this test'))
mockRegisterSSOProvider.mockResolvedValue({ providerId: 'acme-oidc' })
})

afterAll(() => {
resetDbChainMock()
})

it('rejects callers without an Enterprise plan', async () => {
mockHasSSOAccess.mockResolvedValue(false)
const res = await POST(request({ ...OIDC_BODY, orgId: 'org1' }))
Expand All @@ -126,22 +108,22 @@ describe('POST /api/auth/sso/register', () => {
})

it('rejects callers who are not an admin/owner of the target org', async () => {
dbState.members = [{ organizationId: 'org1', role: 'member' }]
queueMembers([{ organizationId: 'org1', role: 'member' }])
const res = await POST(request({ ...OIDC_BODY, orgId: 'org1' }))
expect(res.status).toBe(403)
expect(mockRegisterSSOProvider).not.toHaveBeenCalled()
})

it('rejects an invalid domain', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
const res = await POST(request({ ...OIDC_BODY, domain: 'not-a-domain', orgId: 'org1' }))
expect(res.status).toBe(400)
expect(mockRegisterSSOProvider).not.toHaveBeenCalled()
})

it('rejects a domain already registered by another organization', async () => {
dbState.members = [{ organizationId: 'org-attacker', role: 'owner' }]
dbState.providers = [{ domain: 'acme.com', userId: 'u-victim', organizationId: 'org-victim' }]
queueMembers([{ organizationId: 'org-attacker', role: 'owner' }])
queueProviders([{ domain: 'acme.com', userId: 'u-victim', organizationId: 'org-victim' }])
const res = await POST(request({ ...OIDC_BODY, orgId: 'org-attacker' }))
const json = await res.json()
expect(res.status).toBe(409)
Expand All @@ -150,46 +132,51 @@ describe('POST /api/auth/sso/register', () => {
})

it('matches conflicts across casing variants', async () => {
dbState.members = [{ organizationId: 'org-attacker', role: 'owner' }]
dbState.providers = [{ domain: 'ACME.com', userId: 'u-victim', organizationId: 'org-victim' }]
queueMembers([{ organizationId: 'org-attacker', role: 'owner' }])
queueProviders([{ domain: 'ACME.com', userId: 'u-victim', organizationId: 'org-victim' }])
const res = await POST(request({ ...OIDC_BODY, orgId: 'org-attacker' }))
expect(res.status).toBe(409)
expect(mockRegisterSSOProvider).not.toHaveBeenCalled()
// The conflict lookup itself must be case-insensitive: lower(domain) = <normalized domain>.
const conflictWhere = dbChainMockFns.where.mock.calls.find(([condition]) =>
condition?.strings?.join('?').includes('lower(')
)
expect(conflictWhere?.[0]?.values).toContain('acme.com')
})

it('registers when the domain is unclaimed', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
const res = await POST(request({ ...OIDC_BODY, orgId: 'org1' }))
expect(res.status).toBe(200)
expect(mockRegisterSSOProvider).toHaveBeenCalledTimes(1)
})

it('allows the owning tenant to update its own provider for the same domain', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
dbState.providers = [{ domain: 'acme.com', userId: 'u1', organizationId: 'org1' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
queueProviders([{ domain: 'acme.com', userId: 'u1', organizationId: 'org1' }])
const res = await POST(request({ ...OIDC_BODY, orgId: 'org1' }))
expect(res.status).toBe(200)
expect(mockRegisterSSOProvider).toHaveBeenCalledTimes(1)
})

it('lets an org admin adopt their own user-scoped provider for the same domain', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
dbState.providers = [{ domain: 'acme.com', userId: 'u1', organizationId: null }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
queueProviders([{ domain: 'acme.com', userId: 'u1', organizationId: null }])
const res = await POST(request({ ...OIDC_BODY, orgId: 'org1' }))
expect(res.status).toBe(200)
expect(mockRegisterSSOProvider).toHaveBeenCalledTimes(1)
})

it("still blocks an org admin from claiming another user's user-scoped domain", async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
dbState.providers = [{ domain: 'acme.com', userId: 'someone-else', organizationId: null }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
queueProviders([{ domain: 'acme.com', userId: 'someone-else', organizationId: null }])
const res = await POST(request({ ...OIDC_BODY, orgId: 'org1' }))
expect(res.status).toBe(409)
expect(mockRegisterSSOProvider).not.toHaveBeenCalled()
})

it('normalizes the domain before persisting it', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
const res = await POST(request({ ...OIDC_BODY, domain: 'ACME.com', orgId: 'org1' }))
expect(res.status).toBe(200)
expect(mockRegisterSSOProvider).toHaveBeenCalledTimes(1)
Expand All @@ -198,23 +185,23 @@ describe('POST /api/auth/sso/register', () => {
})

it('passes skipDiscovery since Sim already resolved and validated the OIDC endpoints', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
const res = await POST(request({ ...OIDC_BODY, orgId: 'org1' }))
expect(res.status).toBe(200)
const config = mockRegisterSSOProvider.mock.calls[0][0].body
expect(config.oidcConfig.skipDiscovery).toBe(true)
})

it('omits userInfoEndpoint when skipUserInfoEndpoint is requested, forcing ID token claims', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
const res = await POST(request({ ...OIDC_BODY, skipUserInfoEndpoint: true, orgId: 'org1' }))
expect(res.status).toBe(200)
const config = mockRegisterSSOProvider.mock.calls[0][0].body
expect(config.oidcConfig.userInfoEndpoint).toBeUndefined()
})

it('does not SSRF-validate userInfoEndpoint when skipUserInfoEndpoint is requested', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
mockValidateUrlWithDNS.mockImplementation(async (url: string, label: string) => {
if (label === 'OIDC userInfoEndpoint') {
return { isValid: false, error: 'resolves to a private IP address' }
Expand All @@ -228,7 +215,7 @@ describe('POST /api/auth/sso/register', () => {
})

it('does not SSRF-validate a discovered userinfo_endpoint when skipUserInfoEndpoint is requested', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
mockValidateUrlWithDNS.mockImplementation(async (url: string, label: string) => {
if (label === 'OIDC userinfo_endpoint') {
return { isValid: false, error: 'resolves to a private IP address' }
Expand Down Expand Up @@ -258,15 +245,15 @@ describe('POST /api/auth/sso/register', () => {
})

it('keeps userInfoEndpoint when skipUserInfoEndpoint is not requested', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
const res = await POST(request({ ...OIDC_BODY, orgId: 'org1' }))
expect(res.status).toBe(200)
const config = mockRegisterSSOProvider.mock.calls[0][0].body
expect(config.oidcConfig.userInfoEndpoint).toBe('https://idp.acme.com/userinfo')
})

it('selects tokenEndpointAuthentication from the discovery document when endpoints are auto-discovered', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
mockSecureFetchWithPinnedIP.mockResolvedValue({
ok: true,
json: async () => ({
Expand All @@ -290,7 +277,7 @@ describe('POST /api/auth/sso/register', () => {
})

it('still selects tokenEndpointAuthentication from discovery when all endpoints are explicit', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
mockSecureFetchWithPinnedIP.mockResolvedValue({
ok: true,
json: async () => ({
Expand All @@ -305,7 +292,7 @@ describe('POST /api/auth/sso/register', () => {
})

it('registers successfully when discovery is unreachable and all endpoints are explicit', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
mockSecureFetchWithPinnedIP.mockRejectedValue(new Error('ECONNREFUSED'))
const res = await POST(request({ ...OIDC_BODY, orgId: 'org1' }))
expect(res.status).toBe(200)
Expand All @@ -316,7 +303,7 @@ describe('POST /api/auth/sso/register', () => {
})

it('prefers client_secret_post over client_secret_basic when an IdP supports both', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
mockSecureFetchWithPinnedIP.mockResolvedValue({
ok: true,
json: async () => ({
Expand All @@ -330,7 +317,7 @@ describe('POST /api/auth/sso/register', () => {
})

it('defaults to client_secret_post when discovery advertises no auth methods', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
mockSecureFetchWithPinnedIP.mockResolvedValue({
ok: true,
json: async () => ({}),
Expand All @@ -342,7 +329,7 @@ describe('POST /api/auth/sso/register', () => {
})

it('surfaces the specific discovery failure reason when endpoints are missing', async () => {
dbState.members = [{ organizationId: 'org1', role: 'owner' }]
queueMembers([{ organizationId: 'org1', role: 'owner' }])
mockValidateUrlWithDNS.mockImplementation(async (url: string, label: string) => {
if (label === 'OIDC discovery URL') {
return { isValid: false, error: 'resolves to a private IP address' }
Expand Down
Loading
Loading