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
44 changes: 21 additions & 23 deletions backend/src/middleware/__tests__/loginRateLimit.test.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
import type { NextFunction, Request, Response } from 'express';
import { beforeEach, describe, expect, it, vi } from 'vitest';

import {
clearLoginFailures,
loginRateLimit,
Expand All @@ -7,10 +9,18 @@ import {
} from '../loginRateLimit.js';

function mockRes() {
const res: any = {};
res.status = vi.fn().mockReturnValue(res);
res.json = vi.fn().mockReturnValue(res);
return res;
const res = {
status: vi.fn().mockReturnThis(),
json: vi.fn().mockReturnThis(),
};
return res as unknown as Response & {
status: ReturnType<typeof vi.fn>;
json: ReturnType<typeof vi.fn>;
};
}

function mockReq(ip: string, email: string): Request {
return { ip, body: { email }, socket: {} } as unknown as Request;
}

describe('loginRateLimit', () => {
Expand All @@ -23,28 +33,20 @@ describe('loginRateLimit', () => {
});

it('allows requests under the limit', () => {
const req: any = { ip: '1.1.1.1', body: { email: 'a@x.com' }, socket: {} };
const req = mockReq('1.1.1.1', 'a@x.com');
const res = mockRes();
const next = vi.fn();
const next = vi.fn() as unknown as NextFunction;
loginRateLimit(req, res, next);
expect(next).toHaveBeenCalled();
});

it('429s after per-email max', () => {
const res = mockRes();
for (let i = 0; i < 2; i++) {
loginRateLimit(
{ ip: '2.2.2.2', body: { email: 'b@x.com' }, socket: {} } as any,
mockRes(),
vi.fn()
);
loginRateLimit(mockReq('2.2.2.2', 'b@x.com'), mockRes(), vi.fn() as unknown as NextFunction);
}
const next = vi.fn();
loginRateLimit(
{ ip: '2.2.2.2', body: { email: 'b@x.com' }, socket: {} } as any,
res,
next
);
const next = vi.fn() as unknown as NextFunction;
loginRateLimit(mockReq('2.2.2.2', 'b@x.com'), res, next);
expect(res.status).toHaveBeenCalledWith(429);
expect(next).not.toHaveBeenCalled();
});
Expand All @@ -54,12 +56,8 @@ describe('loginRateLimit', () => {
recordLoginFailure('c@x.com', '3.3.3.3');
recordLoginFailure('c@x.com', '3.3.3.3');
const res = mockRes();
const next = vi.fn();
loginRateLimit(
{ ip: '3.3.3.3', body: { email: 'c@x.com' }, socket: {} } as any,
res,
next
);
const next = vi.fn() as unknown as NextFunction;
loginRateLimit(mockReq('3.3.3.3', 'c@x.com'), res, next);
expect(res.status).toHaveBeenCalledWith(429);
clearLoginFailures('c@x.com', '3.3.3.3');
});
Expand Down
37 changes: 23 additions & 14 deletions backend/src/middleware/__tests__/requireDeploymentAuth.test.ts
Original file line number Diff line number Diff line change
@@ -1,7 +1,10 @@
import { describe, expect, it, vi, beforeEach } from 'vitest';
import type { NextFunction, Response } from 'express';
import { beforeEach, describe, expect, it, vi } from 'vitest';

const getById = vi.fn();
const verifyApiKey = vi.fn();
const { getById, verifyApiKey } = vi.hoisted(() => ({
getById: vi.fn(),
verifyApiKey: vi.fn(),
}));

vi.mock('../../repositories/deploymentRepository.js', () => ({
createDeploymentRepository: () => ({ getById }),
Expand All @@ -20,13 +23,22 @@ vi.mock('../resourceOwnership.js', () => ({
verifyProjectOwnership: vi.fn(),
}));

import type { PredictRequest } from '../requireDeploymentAuth.js';
import { requireDeploymentAuth } from '../requireDeploymentAuth.js';

function mockRes() {
const res: any = {};
res.status = vi.fn().mockReturnValue(res);
res.json = vi.fn().mockReturnValue(res);
return res;
const res = {
status: vi.fn().mockReturnThis(),
json: vi.fn().mockReturnThis(),
};
return res as unknown as Response & {
status: ReturnType<typeof vi.fn>;
json: ReturnType<typeof vi.fn>;
};
}

function mockReq(params: Record<string, string>, headers: Record<string, string> = {}): PredictRequest {
return { params, headers } as unknown as PredictRequest;
}

describe('requireDeploymentAuth enumeration resistance', () => {
Expand All @@ -37,9 +49,9 @@ describe('requireDeploymentAuth enumeration resistance', () => {

it('returns 404 Not found for missing deployment', async () => {
getById.mockResolvedValue(null);
const req: any = { params: { deploymentId: 'missing' }, headers: {} };
const req = mockReq({ deploymentId: 'missing' });
const res = mockRes();
const next = vi.fn();
const next = vi.fn() as unknown as NextFunction;
await requireDeploymentAuth(req, res, next);
expect(res.status).toHaveBeenCalledWith(404);
expect(res.json).toHaveBeenCalledWith({ error: 'Not found' });
Expand All @@ -49,12 +61,9 @@ describe('requireDeploymentAuth enumeration resistance', () => {
it('returns the same 404 Not found for a bad API key', async () => {
getById.mockResolvedValue({ deploymentId: 'dep-1', projectId: 'p1' });
verifyApiKey.mockResolvedValue(null);
const req: any = {
params: { deploymentId: 'dep-1' },
headers: { 'x-api-key': 'bad-key' },
};
const req = mockReq({ deploymentId: 'dep-1' }, { 'x-api-key': 'bad-key' });
const res = mockRes();
const next = vi.fn();
const next = vi.fn() as unknown as NextFunction;
await requireDeploymentAuth(req, res, next);
expect(res.status).toHaveBeenCalledWith(404);
expect(res.json).toHaveBeenCalledWith({ error: 'Not found' });
Expand Down
2 changes: 1 addition & 1 deletion backend/src/routes/auth.ts
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,8 @@ import { z } from 'zod';
import { appLogger } from '../logging/logger.js';
import { asyncHandler } from '../middleware/asyncHandler.js';
import { requireAuth, requireAuthAllowUnverified, invalidateUserCache } from '../middleware/auth.js';
import { validateRequest } from '../middleware/validateRequest.js';
import { loginRateLimit, recordLoginFailure, clearLoginFailures } from '../middleware/loginRateLimit.js';
import { validateRequest } from '../middleware/validateRequest.js';
import { UserRepository } from '../repositories/userRepository.js';
import { authService } from '../services/authService.js';
import { emailService } from '../services/emailService.js';
Expand Down
82 changes: 57 additions & 25 deletions backend/src/routes/auth/__tests__/oauthHandler.test.ts
Original file line number Diff line number Diff line change
@@ -1,5 +1,20 @@
import type { Request, Response } from 'express';
import { beforeEach, describe, expect, it, vi } from 'vitest';

const {
hashPassword,
generatePasswordResetToken,
generateTokens,
hashRefreshToken,
refreshTokenExpiryMs,
} = vi.hoisted(() => ({
hashPassword: vi.fn().mockResolvedValue('hash'),
generatePasswordResetToken: vi.fn().mockReturnValue('rand'),
generateTokens: vi.fn().mockReturnValue({ accessToken: 'a', refreshToken: 'r' }),
hashRefreshToken: vi.fn().mockReturnValue('rh'),
refreshTokenExpiryMs: vi.fn().mockReturnValue(1000),
}));

vi.mock('../../../config.js', () => ({
env: {
googleClientId: 'cid',
Expand All @@ -12,12 +27,6 @@ vi.mock('../../../logging/logger.js', () => ({
appLogger: { error: vi.fn(), warn: vi.fn(), info: vi.fn() },
}));

const hashPassword = vi.fn().mockResolvedValue('hash');
const generatePasswordResetToken = vi.fn().mockReturnValue('rand');
const generateTokens = vi.fn().mockReturnValue({ accessToken: 'a', refreshToken: 'r' });
const hashRefreshToken = vi.fn().mockReturnValue('rh');
const refreshTokenExpiryMs = vi.fn().mockReturnValue(1000);

vi.mock('../../../services/authService.js', () => ({
authService: {
hashPassword,
Expand All @@ -28,58 +37,81 @@ vi.mock('../../../services/authService.js', () => ({
},
}));

import type { UserRepository } from '../../../repositories/userRepository.js';
import { handleGoogleCallback } from '../oauthHandler.js';

type GoogleUser = {
id: string;
email: string;
name: string;
verified_email: boolean;
};

const googleUserStore: { current: GoogleUser | null } = { current: null };

function mockRes() {
const res: any = {};
res.status = vi.fn().mockReturnValue(res);
res.json = vi.fn().mockReturnValue(res);
return res;
const res = {
status: vi.fn().mockReturnThis(),
json: vi.fn().mockReturnThis(),
};
return res as unknown as Response & {
status: ReturnType<typeof vi.fn>;
json: ReturnType<typeof vi.fn>;
};
}

function mockRepo(overrides: Record<string, unknown> = {}) {
function mockRepo(overrides: Partial<Record<keyof UserRepository, unknown>> = {}): UserRepository {
return {
findByEmail: vi.fn().mockResolvedValue(null),
create: vi.fn(),
markEmailVerified: vi.fn(),
findById: vi.fn(),
updateLastLogin: vi.fn(),
toSafeUser: vi.fn((u) => u),
toSafeUser: vi.fn((u: unknown) => u),
storeRefreshToken: vi.fn(),
...overrides,
} as any;
} as unknown as UserRepository;
}

function mockReq(): Request {
return { body: { code: 'c' }, ip: '1', get: () => 'ua' } as unknown as Request;
}

function mockFetchResponse(body: unknown): globalThis.Response {
return {
ok: true,
json: async () => body,
} as unknown as globalThis.Response;
}

describe('handleGoogleCallback', () => {
beforeEach(() => {
googleUserStore.current = null;
vi.stubGlobal(
'fetch',
vi.fn(async (url: string) => {
vi.fn(async (url: string | URL | globalThis.Request) => {
if (String(url).includes('oauth2.googleapis.com/token')) {
return { ok: true, json: async () => ({ access_token: 'tok', id_token: 'id' }) } as any;
return mockFetchResponse({ access_token: 'tok', id_token: 'id' });
}
return {
ok: true,
json: async () => (globalThis as any).__googleUser,
} as any;
return mockFetchResponse(googleUserStore.current);
})
);
});

it('rejects unverified Google emails', async () => {
(globalThis as any).__googleUser = {
googleUserStore.current = {
id: 'g1',
email: 'a@x.com',
name: 'A',
verified_email: false,
};
const res = mockRes();
await handleGoogleCallback({ body: { code: 'c' }, ip: '1', get: () => 'ua' } as any, res, mockRepo());
await handleGoogleCallback(mockReq(), res, mockRepo());
expect(res.status).toHaveBeenCalledWith(403);
});

it('blocks silent merge into unverified password account', async () => {
(globalThis as any).__googleUser = {
googleUserStore.current = {
id: 'g1',
email: 'a@x.com',
name: 'A',
Expand All @@ -88,15 +120,15 @@ describe('handleGoogleCallback', () => {
const existing = { user_id: 'u1', email: 'a@x.com', email_verified: false };
const res = mockRes();
await handleGoogleCallback(
{ body: { code: 'c' }, ip: '1', get: () => 'ua' } as any,
mockReq(),
res,
mockRepo({ findByEmail: vi.fn().mockResolvedValue(existing) })
);
expect(res.status).toHaveBeenCalledWith(409);
});

it('logs into verified existing account', async () => {
(globalThis as any).__googleUser = {
googleUserStore.current = {
id: 'g1',
email: 'a@x.com',
name: 'A',
Expand All @@ -108,7 +140,7 @@ describe('handleGoogleCallback', () => {
toSafeUser: vi.fn().mockReturnValue(existing),
});
const res = mockRes();
await handleGoogleCallback({ body: { code: 'c' }, ip: '1', get: () => 'ua' } as any, res, repo);
await handleGoogleCallback(mockReq(), res, repo);
expect(repo.updateLastLogin).toHaveBeenCalledWith('u1');
expect(res.json).toHaveBeenCalled();
expect(res.status).not.toHaveBeenCalledWith(409);
Expand Down
2 changes: 1 addition & 1 deletion backend/src/routes/datasets.ts
Original file line number Diff line number Diff line change
Expand Up @@ -13,8 +13,8 @@ import { getProjectRepository } from '../repositories/projectRepository.js';
import { resolveDatasetTableName } from '../services/datasetLoader.js';
import { ensureProjectDatasetSqlNames, resolveDatasetSqlName } from '../services/datasetSqlNames.js';
import { getDatasetQueryState, rebuildDatasetTableFromSource } from '../services/datasetTableManager.js';
import { getWorkflowRepository } from '../services/workflows/repository/index.js';
import * as notebookService from '../services/notebook/notebookService.js';
import { getWorkflowRepository } from '../services/workflows/repository/index.js';
import type { AuthRequest } from '../types/auth.js';
import { getErrorMessage, sendNotFound } from '../utils/errors.js';
import { getDatasetPath } from '../utils/pathUtils.js';
Expand Down
2 changes: 1 addition & 1 deletion backend/src/routes/notebooks/notebookRoutes.ts
Original file line number Diff line number Diff line change
Expand Up @@ -6,11 +6,11 @@ import { verifyProjectOwnership } from '../../middleware/resourceOwnership.js';
import { getProjectRepository } from '../../repositories/projectRepository.js';
import * as kernelManager from '../../services/kernelManager.js';
import { getOrEnsureContainer } from '../../services/notebook/cellExecutionService.js';
import * as notebookService from '../../services/notebook/notebookService.js';
import {
getNotebookRecoveryCandidate,
recoverNotebookFromWorkflowHistory
} from '../../services/notebook/notebookRecoveryService.js';
import * as notebookService from '../../services/notebook/notebookService.js';
import type { AuthRequest } from '../../types/auth.js';
import { NotebookKindSchema } from '../../types/notebook.js';

Expand Down
2 changes: 1 addition & 1 deletion backend/src/services/llm/prompts/trainingWorkflow.ts
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,8 @@

import type { DatasetProfile } from '../../../types/dataset.js';
import type { ToolResult } from '../../../types/llm.js';
import type { FeatureSpec } from '../../featureEngineering.js';
import { findLikelyIdentifierColumns } from '../../columnClassification.js';
import type { FeatureSpec } from '../../featureEngineering.js';
import { buildTemplateSummary } from '../../modelTemplates.js';
import type {
LlmRequest,
Expand Down
3 changes: 2 additions & 1 deletion backend/src/services/llm/trainingTools/executionTools.ts
Original file line number Diff line number Diff line change
@@ -1,10 +1,11 @@
import { appLogger } from '../../../logging/logger.js';
import { executeMcpTool } from '../../mcp/mcpAdapter.js';
import { nowIso } from '../preprocessingTools/helpers.js';
import { normalizeWorkflowPrepSegments } from './workflowPrepSegments.js';


import { resolveExperiment } from './types.js';
import type { TrainingToolContext, TrainingToolHandler, TrainingToolResult } from './types.js';
import { normalizeWorkflowPrepSegments } from './workflowPrepSegments.js';

function normalizeMetricsRecord(metrics: unknown): Record<string, number> {
if (!metrics || typeof metrics !== 'object' || Array.isArray(metrics)) {
Expand Down
2 changes: 1 addition & 1 deletion backend/src/services/modelTraining.ts
Original file line number Diff line number Diff line change
Expand Up @@ -15,8 +15,8 @@ import { syncWorkspaceDatasets } from './executionWorkspace.js';
import * as kernelManager from './kernelManager.js';
import { getModelTemplate, listModelTemplates, resolveModelTemplateId } from './modelTemplates.js';
import { DEFAULT_MODEL_TEST_SIZE, normalizeModelTestSize } from './modelTestSize.js';
import { deleteTuningStudiesByModelId } from './tuningService.js';
import { resolveContainerSafeNJobs } from './parallelism.js';
import { deleteTuningStudiesByModelId } from './tuningService.js';

const datasetRepository = createDatasetRepository(env.datasetMetadataPath);
const modelRepository = createModelRepository(env.modelMetadataPath);
Expand Down
5 changes: 3 additions & 2 deletions backend/src/services/notebook/notebookRecoveryService.ts
Original file line number Diff line number Diff line change
@@ -1,8 +1,9 @@
import * as notebookRepo from '../../repositories/notebookRepository.js';
import type { CellOutput, CellStatus } from '../../types/notebook.js';
import type { WorkflowPhase, WorkflowRunSnapshot } from '../workflows/types.js';
import { getWorkflowRepository } from '../workflows/repository/index.js';
import type { WorkflowPhase, WorkflowRunSnapshot } from '../workflows/types.js';

import * as notebookService from './notebookService.js';
import * as notebookRepo from '../../repositories/notebookRepository.js';

type RecoverablePhase = 'preprocessing' | 'feature-engineering' | 'training';

Expand Down
1 change: 1 addition & 0 deletions backend/src/services/parallelism.test.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import { describe, expect, it } from 'vitest';

import { buildNJobsPythonSnippet, resolveContainerSafeNJobs } from './parallelism.js';

describe('resolveContainerSafeNJobs', () => {
Expand Down
2 changes: 1 addition & 1 deletion backend/src/services/tuningService.ts
Original file line number Diff line number Diff line change
Expand Up @@ -26,8 +26,8 @@ import { resolveAndHealTargetColumn } from '../utils/modelUtils.js';

import { getModelTemplate } from './modelTemplates.js';
import { resolveModelTestSize } from './modelTestSize.js';
import {
import { buildNJobsPythonSnippet } from './parallelism.js';
import {
buildOutputDirSetup,
buildResultSaving,
buildStandardImports,
Expand Down
Loading
Loading