Skip to content
Open
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
106 changes: 106 additions & 0 deletions src/allowlist/index.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,106 @@
import { describe, it, expect, vi, beforeEach } from 'vitest'
import { isQueryAllowed } from './index'
import type { DataSource } from '../types'
import type { StarbaseDBConfiguration } from '../handler'

describe('Allowlist Module', () => {
let mockDataSource: DataSource
let mockConfig: StarbaseDBConfiguration

beforeEach(() => {
vi.clearAllMocks()

mockDataSource = {
source: 'internal',
cache: false,
rpc: {
executeQuery: vi.fn(),
},
} as any

mockConfig = {
role: 'client',
features: {
allowlist: true,
rls: false,
},
} as any
})

it('should return true if feature is disabled', async () => {
const result = await isQueryAllowed({
sql: 'SELECT * FROM users',
isEnabled: false,
dataSource: mockDataSource,
config: mockConfig,
})
expect(result).toBe(true)
})

it('should return true if user is admin', async () => {
mockConfig.role = 'admin'
const result = await isQueryAllowed({
sql: 'SELECT * FROM users',
isEnabled: true,
dataSource: mockDataSource,
config: mockConfig,
})
expect(result).toBe(true)
})

it('should return error if no SQL is provided', async () => {
const result = await isQueryAllowed({
sql: '',
isEnabled: true,
dataSource: mockDataSource,
config: mockConfig,
})
expect(result).toBeInstanceOf(Error)
expect((result as Error).message).toBe(
'No SQL provided for allowlist check'
)
})

it('should return true if query is in allowlist', async () => {
vi.mocked(mockDataSource.rpc.executeQuery).mockResolvedValue([
{ sql_statement: 'SELECT * FROM users', source: 'internal' },
])

const result = await isQueryAllowed({
sql: 'SELECT * FROM users',
isEnabled: true,
dataSource: mockDataSource,
config: mockConfig,
})
expect(result).toBe(true)
})

it('should throw if query is not in allowlist', async () => {
vi.mocked(mockDataSource.rpc.executeQuery).mockResolvedValue([
{ sql_statement: 'SELECT * FROM users', source: 'internal' },
])

await expect(
isQueryAllowed({
sql: 'SELECT * FROM admins',
isEnabled: true,
dataSource: mockDataSource,
config: mockConfig,
})
).rejects.toThrow('Query not allowed')
})

it('should allow queries with different whitespace/semicolons but same AST', async () => {
vi.mocked(mockDataSource.rpc.executeQuery).mockResolvedValue([
{ sql_statement: 'SELECT * FROM users', source: 'internal' },
])

const result = await isQueryAllowed({
sql: 'SELECT * \nFROM \t users; ',
isEnabled: true,
dataSource: mockDataSource,
config: mockConfig,
})
expect(result).toBe(true)
})
})
121 changes: 121 additions & 0 deletions src/do.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -145,4 +145,125 @@ describe('StarbaseDBDurableObject Tests', () => {
instance.executeQuery({ sql: 'INVALID QUERY' })
).rejects.toThrow('Query failed')
})

it('should handle getAlarm, setAlarm, deleteAlarm', async () => {
mockDurableObjectState.storage.getAlarm = vi
.fn()
.mockResolvedValue(12345)
mockDurableObjectState.storage.setAlarm = vi
.fn()
.mockResolvedValue(undefined)
mockDurableObjectState.storage.deleteAlarm = vi
.fn()
.mockResolvedValue(undefined)

expect(await instance.getAlarm()).toBe(12345)

await instance.setAlarm(new Date(Date.now() + 5000))
expect(mockDurableObjectState.storage.setAlarm).toHaveBeenCalled()

await instance.deleteAlarm()
expect(mockDurableObjectState.storage.deleteAlarm).toHaveBeenCalled()
})

it('should handle alarm execution', async () => {
mockStorage.sql.exec.mockReturnValueOnce({
toArray: vi.fn().mockReturnValue([
{
is_active: 1,
callback_host: 'http://localhost',
},
]),
})
global.fetch = vi.fn().mockResolvedValue({ status: 200 } as any)

await instance.alarm()
expect(global.fetch).toHaveBeenCalled()
})

it('should handle alarm execution when task fails', async () => {
mockStorage.sql.exec.mockReturnValueOnce({
toArray: vi.fn().mockReturnValue([
{
is_active: 1,
callback_host: 'http://localhost',
},
]),
})
global.fetch = vi.fn().mockRejectedValue(new Error('Network error'))
mockDurableObjectState.storage.setAlarm = vi.fn()

await instance.alarm()
expect(mockDurableObjectState.storage.setAlarm).toHaveBeenCalled() // Recovery alarm
})

it('should return statistics', async () => {
instance.sql.databaseSize = 1024
mockStorage.sql.exec.mockReturnValueOnce({
toArray: vi.fn().mockReturnValue([{ count: 5 }]),
})
const stats = await instance.getStatistics()
expect(stats.databaseSize).toBe(1024)
expect(stats.activeConnections).toBe(0)
expect(stats.recentQueries).toBe(5)
})

it('should reject /socket if not websocket upgrade', async () => {
const request = new Request('http://localhost/socket')
const response = await instance.fetch(request)
expect(response.status).toBe(400)
})

it('should handle /socket with upgrade', async () => {
const request = new Request('http://localhost/socket?sessionId=test', {
headers: { upgrade: 'websocket' },
})
const response = await instance.fetch(request)
expect(response.status).toBe(101)
expect(instance.connections.has('test')).toBe(true)
})

it('should handle /socket/broadcast', async () => {
const server = new global.WebSocket('ws://localhost')
server.send = vi.fn()
instance.connections.set('test1', server)

const request = new Request('http://localhost/socket/broadcast', {
method: 'POST',
body: JSON.stringify({ msg: 'hello' }),
})
const response = await instance.fetch(request)
expect(response.status).toBe(200)
expect(server.send).toHaveBeenCalled()
})

it('should handle webSocketMessage', async () => {
const ws = new global.WebSocket('ws://localhost')
ws.send = vi.fn()
await instance.webSocketMessage(
ws,
JSON.stringify({ action: 'query', sql: 'SELECT 1' })
)
expect(ws.send).toHaveBeenCalled()
})

it('should handle webSocketClose', async () => {
const ws = new global.WebSocket('ws://localhost')
ws.close = vi.fn()
instance.ctx = mockDurableObjectState
instance.connections.set('session-123', ws)
await instance.webSocketClose(ws, 1000, 'done', true)
expect(ws.close).toHaveBeenCalled()
expect(instance.connections.has('session-123')).toBe(false)
})

it('should execute raw queries', async () => {
const result = await instance.executeQuery({
sql: 'SELECT 1',
isRaw: true,
params: [1],
})
expect(result).toHaveProperty('columns')
expect(result).toHaveProperty('rows')
})
})
135 changes: 135 additions & 0 deletions src/handler.extra.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,135 @@
import { describe, it, expect, vi, beforeEach } from 'vitest'
import { StarbaseDB } from './handler'
import type { DataSource } from './types'

describe('handler additional tests', () => {
let mockDataSource: DataSource
let mockConfig: any

beforeEach(() => {
mockConfig = {
role: 'client',
features: {},
}

mockDataSource = {
source: 'internal',
rpc: {
executeQuery: vi.fn(),
} as any,
} as DataSource
})

it('should throw if external data source is missing external config', () => {
mockDataSource.source = 'external'
expect(() => {
new StarbaseDB({
dataSource: mockDataSource,
config: mockConfig,
})
}).toThrow('No external data sources available.')
})

it('should handle preAuth with authless plugin', async () => {
const instance = new StarbaseDB({
dataSource: mockDataSource,
config: mockConfig,
plugins: [
{
opts: { requiresAuth: false },
pathPrefix: '/public/.*',
register: vi.fn(),
beforeQuery: vi.fn(),
afterQuery: vi.fn(),
} as any,
],
})
const req = new Request('http://localhost/public/hello')
const res = await instance.handlePreAuth(req, {} as any)
expect(res).toBeDefined()
})

it('should handle preAuth returning undefined when auth required', async () => {
const instance = new StarbaseDB({
dataSource: mockDataSource,
config: mockConfig,
plugins: [
{
opts: { requiresAuth: true },
pathPrefix: '/private/*',
register: vi.fn(),
beforeQuery: vi.fn(),
afterQuery: vi.fn(),
} as any,
],
})
const req = new Request('http://localhost/private/hello')
const res = await instance.handlePreAuth(req, {} as any)
expect(res).toBeUndefined()
})

it('should reject non-json in queryRoute', async () => {
const instance = new StarbaseDB({
dataSource: mockDataSource,
config: mockConfig,
})
const req = new Request('http://localhost', {
method: 'POST',
headers: { 'Content-Type': 'text/plain' },
body: 'hello',
})
const res = await instance.queryRoute(req, false)
expect(res.status).toBe(400)
})

it('should reject invalid sql field in transaction', async () => {
const instance = new StarbaseDB({
dataSource: mockDataSource,
config: mockConfig,
})
const req = new Request('http://localhost', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ transaction: [{ sql: 123 }] }),
})
const res = await instance.queryRoute(req, false)
expect(res.status).toBe(500)
expect((await res.json()).error).toContain(
'Invalid or empty "sql" field in transaction'
)
})

it('should reject invalid params field in transaction', async () => {
const instance = new StarbaseDB({
dataSource: mockDataSource,
config: mockConfig,
})
const req = new Request('http://localhost', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({
transaction: [{ sql: 'SELECT 1', params: 'invalid' }],
}),
})
const res = await instance.queryRoute(req, false)
expect(res.status).toBe(500)
expect((await res.json()).error).toContain(
'Invalid "params" field in transaction'
)
})

it('should reject invalid params field in queryRoute', async () => {
const instance = new StarbaseDB({
dataSource: mockDataSource,
config: mockConfig,
})
const req = new Request('http://localhost', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ sql: 'SELECT 1', params: 'invalid' }),
})
const res = await instance.queryRoute(req, false)
expect(res.status).toBe(400)
expect((await res.json()).error).toContain('Invalid "params" field')
})
})
Loading