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
2 changes: 2 additions & 0 deletions src/services/api/xaiOAuth.ts
Original file line number Diff line number Diff line change
Expand Up @@ -287,6 +287,7 @@ export class XaiOAuthService {
* The loopback callback server stays open until tokens are obtained, the
* caller submits a manual code, or `cancel()` is invoked. CORS preflight
* from `auth.x.ai` is echoed so xAI's browser-side fetch can reach us.
* Callbacks with missing or mismatched state leave the login pending.
*/
async beginOAuthFlow(): Promise<XaiOAuthFlowHandle> {
// Reset cross-flow state so a reused service instance starts clean.
Expand All @@ -310,6 +311,7 @@ export class XaiOAuthService {
port: callbackPort,
host: callbackHost,
callbackPath: XAI_OAUTH_CALLBACK_PATH,
expectedState: state,
successTitle: 'xAI OAuth complete',
})
} catch (error) {
Expand Down
124 changes: 121 additions & 3 deletions src/services/api/xaiOAuthCallback.test.ts
Original file line number Diff line number Diff line change
@@ -1,14 +1,16 @@
import { afterEach, beforeEach, describe, expect, test } from 'bun:test'
import { afterEach, beforeEach, describe, expect, spyOn, test } from 'bun:test'
import { connect } from 'node:net'

import { acquireSharedMutationLock, releaseSharedMutationLock } from '../../test/sharedMutationLock.js'
import { startXaiOAuthCallback } from './xaiOAuthCallback.js'
import { XaiOAuthService } from './xaiOAuth.js'

async function startTestServer() {
const handle = await startXaiOAuthCallback({
port: 0,
host: '127.0.0.1',
callbackPath: '/callback',
expectedState: 'xyz',
successTitle: 'xAI OAuth complete',
})
return { handle, port: handle.port }
Expand Down Expand Up @@ -235,16 +237,131 @@ describe.serial('startXaiOAuthCallback (CORS-aware loopback for xAI auth)', () =
expect(result).toEqual({ code: 'ABC123', state: 'xyz' })
})

test('GET with ?error=access_denied rejects with a clear message', async () => {
test('GET with a matching state and OAuth error rejects with a clear message', async () => {
const { handle, port } = await startTestServer()
cleanup = () => handle.close()

const callbackPromise = handle.waitForCallback()
const res = await requestLoopback(port, '/callback?error=access_denied')
const res = await requestLoopback(port, '/callback?error=access_denied&state=xyz')
expect(res.status).toBe(400)
await expect(callbackPromise).rejects.toThrow(/access_denied/)
})

for (const query of [
'error=access_denied',
'error=access_denied&state=wrong',
'error=access_denied&state=',
'code=forged&state=wrong',
'code=forged',
'state=wrong',
'',
'code=forged&state=%20xyz%20',
]) {
test(`invalid state does not consume the callback: ${query || '(empty query)'}`, async () => {
const { handle, port } = await startTestServer()
cleanup = () => handle.close()
let settled = false
const callbackPromise = handle.waitForCallback()
void callbackPromise.then(
() => {
settled = true
},
() => {
settled = true
},
)

const rejected = await requestLoopback(port, `/callback?${query}`)
expect(rejected.status).toBe(400)
expect(settled).toBe(false)

const accepted = await requestLoopback(
port,
'/callback?code=legitimate&state=xyz',
)
expect(accepted.status).toBe(200)
await expect(callbackPromise).resolves.toEqual({
code: 'legitimate',
state: 'xyz',
})
})
}

for (const completion of ['callback', 'manual', 'cancel'] as const) {
test(`OAuth service survives an invalid request before ${completion}`, async () => {
const exchangedCodes: string[] = []
const fetchSpy = spyOn(globalThis, 'fetch').mockImplementation(
Object.assign(
async (input: Parameters<typeof fetch>[0], init?: RequestInit) => {
const url = String(input)
if (url.endsWith('/.well-known/openid-configuration')) {
return Response.json({
authorization_endpoint: 'https://auth.x.ai/authorize',
token_endpoint: 'https://auth.x.ai/token',
})
}
expect(url).toBe('https://auth.x.ai/token')
const body = new URLSearchParams(String(init?.body))
exchangedCodes.push(body.get('code') ?? '')
return Response.json({
access_token: 'test-access',
refresh_token: 'test-refresh',
})
},
{ preconnect: globalThis.fetch.preconnect },
),
)
const service = new XaiOAuthService({
callbackPort: 0,
callbackHost: '127.0.0.1',
})
try {
const flow = await service.beginOAuthFlow()
const authUrl = new URL(flow.authUrl)
const state = authUrl.searchParams.get('state')!
const redirect = new URL(authUrl.searchParams.get('redirect_uri')!)
const pending = flow.waitForTokens()
let settled = false
void pending.then(
() => {
settled = true
},
() => {
settled = true
},
)
const invalid = await requestLoopback(
Number(redirect.port),
'/callback?error=access_denied',
)
expect(invalid.status).toBe(400)
expect(settled).toBe(false)
expect(exchangedCodes).toEqual([])

if (completion === 'cancel') {
flow.cancel()
await expect(pending).rejects.toThrow(/cancelled|closed/)
expect(exchangedCodes).toEqual([])
} else {
if (completion === 'callback') {
const res = await requestLoopback(
Number(redirect.port),
`/callback?code=legitimate&state=${encodeURIComponent(state)}`,
)
expect(res.status).toBe(200)
} else {
flow.submitManualCode('legitimate')
}
expect((await pending).accessToken).toBe('test-access')
expect(exchangedCodes).toEqual(['legitimate'])
}
} finally {
service.cleanup()
fetchSpy.mockRestore()
}
})
}

test('GET to wrong path returns 404 and does not settle the callback', async () => {
const { handle, port } = await startTestServer()
cleanup = () => handle.close()
Expand Down Expand Up @@ -282,6 +399,7 @@ describe.serial('startXaiOAuthCallback (CORS-aware loopback for xAI auth)', () =
port: 0,
host: '127.0.0.1',
callbackPath: '/callback',
expectedState: 'B',
successTitle: '<script>alert(1)</script>',
})
cleanup = () => handle.close()
Expand Down
14 changes: 12 additions & 2 deletions src/services/api/xaiOAuthCallback.ts
Original file line number Diff line number Diff line change
Expand Up @@ -90,6 +90,7 @@ export async function startXaiOAuthCallback(params: {
port: number
host: string
callbackPath: string
expectedState: string
successTitle?: string
corsOriginAllowlist?: readonly string[]
}): Promise<XaiOAuthCallbackHandle> {
Expand Down Expand Up @@ -153,6 +154,16 @@ export async function startXaiOAuthCallback(params: {
return
}

// Unauthenticated requests must not consume the pending OAuth flow.
// Check the exact state before either error or code can settle it.
const state = url.searchParams.get('state')
if (!state || state !== params.expectedState) {
res.statusCode = 400
res.setHeader('Content-Type', 'text/plain')
res.end('Invalid state')
return
}

const error = url.searchParams.get('error')
if (error) {
res.statusCode = 400
Expand All @@ -166,9 +177,8 @@ export async function startXaiOAuthCallback(params: {
}

const code = url.searchParams.get('code')?.trim()
const state = url.searchParams.get('state')?.trim()

if (!code || !state) {
if (!code) {
res.statusCode = 400
res.setHeader('Content-Type', 'text/plain')
res.end('Missing code or state')
Expand Down
Loading