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
47 changes: 47 additions & 0 deletions packages/core/src/code_assist/codeAssist.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -15,18 +15,23 @@ import {
} from './codeAssist.js';
import type { Config } from '../config/config.js';
import { LoggingContentGenerator } from '../core/loggingContentGenerator.js';
import { ModelMappingContentGenerator } from '../core/modelMappingContentGenerator.js';
import { UserTierId } from './types.js';

// Mock dependencies
vi.mock('./oauth2.js');
vi.mock('./setup.js');
vi.mock('./server.js');
vi.mock('../core/loggingContentGenerator.js');
vi.mock('../core/modelMappingContentGenerator.js');

const mockedGetOauthClient = vi.mocked(getOauthClient);
const mockedSetupUser = vi.mocked(setupUser);
const MockedCodeAssistServer = vi.mocked(CodeAssistServer);
const MockedLoggingContentGenerator = vi.mocked(LoggingContentGenerator);
const MockedModelMappingContentGenerator = vi.mocked(
ModelMappingContentGenerator,
);

describe('codeAssist', () => {
beforeEach(() => {
Expand Down Expand Up @@ -178,5 +183,47 @@ describe('codeAssist', () => {
const server = getCodeAssistServer(mockConfig);
expect(server).toBeUndefined();
});

it('should unwrap and return the server if it is wrapped in a ModelMappingContentGenerator', () => {
const mockServer = new MockedCodeAssistServer({} as never, '', {});
const mockMapper = new MockedModelMappingContentGenerator(
{} as never,
{},
);
vi.spyOn(mockMapper, 'getWrapped').mockReturnValue(mockServer);

const mockConfig = {
getContentGenerator: () => mockMapper,
} as unknown as Config;

const server = getCodeAssistServer(mockConfig);
expect(server).toBe(mockServer);
expect(mockMapper.getWrapped).toHaveBeenCalled();
});

it('should recursively unwrap multiple layers of LoggingContentGenerator and ModelMappingContentGenerator', () => {
const mockServer = new MockedCodeAssistServer({} as never, '', {});
const mockLogger = new MockedLoggingContentGenerator(
{} as never,
{} as never,
);
const mockMapper = new MockedModelMappingContentGenerator(
{} as never,
{},
);

// Mapper wraps Logger wraps Server
vi.spyOn(mockMapper, 'getWrapped').mockReturnValue(mockLogger);
vi.spyOn(mockLogger, 'getWrapped').mockReturnValue(mockServer);

const mockConfig = {
getContentGenerator: () => mockMapper,
} as unknown as Config;

const server = getCodeAssistServer(mockConfig);
expect(server).toBe(mockServer);
expect(mockMapper.getWrapped).toHaveBeenCalled();
expect(mockLogger.getWrapped).toHaveBeenCalled();
});
});
});
13 changes: 10 additions & 3 deletions packages/core/src/code_assist/codeAssist.ts
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import { setupUser } from './setup.js';
import { CodeAssistServer, type HttpOptions } from './server.js';
import type { Config } from '../config/config.js';
import { LoggingContentGenerator } from '../core/loggingContentGenerator.js';
import { ModelMappingContentGenerator } from '../core/modelMappingContentGenerator.js';

export async function createCodeAssistContentGenerator(
httpOptions: HttpOptions,
Expand Down Expand Up @@ -43,9 +44,15 @@ export function getCodeAssistServer(
): CodeAssistServer | undefined {
let server = config.getContentGenerator();

// Unwrap LoggingContentGenerator if present
if (server instanceof LoggingContentGenerator) {
server = server.getWrapped();
// Recursively unwrap LoggingContentGenerator and ModelMappingContentGenerator
while (true) {
if (server instanceof LoggingContentGenerator) {
server = server.getWrapped();
} else if (server instanceof ModelMappingContentGenerator) {
server = server.getWrapped();
} else {
break;
}
}

if (!(server instanceof CodeAssistServer)) {
Expand Down
6 changes: 3 additions & 3 deletions packages/core/src/config/config.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -4357,7 +4357,7 @@
cwd: '.',
};

it('should set DEFAULT_GEMINI_FLASH_MODEL to gemini-3.5-flash and PREVIEW_GEMINI_FLASH_MODEL to gemini-3-flash-preview if hasGemini35FlashGAAccess returns true and authType is USE_GEMINI', () => {

Check warning on line 4360 in packages/core/src/config/config.test.ts

View workflow job for this annotation

GitHub Actions / Lint

Found sensitive keyword "gemini-3.5". Please make sure this change is appropriate to submit.
const config = new Config(baseParams);
config['contentGeneratorConfig'] = { authType: AuthType.USE_GEMINI };

Expand All @@ -4375,11 +4375,11 @@
const result = config.hasGemini35FlashGAAccess();
expect(result).toBe(true);

expect(DEFAULT_GEMINI_FLASH_MODEL).toBe('gemini-3.5-flash');

Check warning on line 4378 in packages/core/src/config/config.test.ts

View workflow job for this annotation

GitHub Actions / Lint

Found sensitive keyword "gemini-3.5". Please make sure this change is appropriate to submit.
expect(PREVIEW_GEMINI_FLASH_MODEL).toBe('gemini-3-flash-preview');
});

it('should set DEFAULT_GEMINI_FLASH_MODEL and PREVIEW_GEMINI_FLASH_MODEL to gemini-3-flash if hasGemini35FlashGAAccess returns true and authType is not USE_GEMINI', () => {
it('should set DEFAULT_GEMINI_FLASH_MODEL and PREVIEW_GEMINI_FLASH_MODEL to gemini-3.5-flash if hasGemini35FlashGAAccess returns true and authType is not USE_GEMINI', () => {

Check warning on line 4382 in packages/core/src/config/config.test.ts

View workflow job for this annotation

GitHub Actions / Lint

Found sensitive keyword "gemini-3.5". Please make sure this change is appropriate to submit.
const config = new Config(baseParams);
config['contentGeneratorConfig'] = { authType: AuthType.LOGIN_WITH_GOOGLE };

Expand All @@ -4397,7 +4397,7 @@
const result = config.hasGemini35FlashGAAccess();
expect(result).toBe(true);

expect(DEFAULT_GEMINI_FLASH_MODEL).toBe('gemini-3-flash');
expect(PREVIEW_GEMINI_FLASH_MODEL).toBe('gemini-3-flash');
expect(DEFAULT_GEMINI_FLASH_MODEL).toBe('gemini-3.5-flash');

Check warning on line 4400 in packages/core/src/config/config.test.ts

View workflow job for this annotation

GitHub Actions / Lint

Found sensitive keyword "gemini-3.5". Please make sure this change is appropriate to submit.
expect(PREVIEW_GEMINI_FLASH_MODEL).toBe('gemini-3.5-flash');

Check warning on line 4401 in packages/core/src/config/config.test.ts

View workflow job for this annotation

GitHub Actions / Lint

Found sensitive keyword "gemini-3.5". Please make sure this change is appropriate to submit.
});
});
2 changes: 1 addition & 1 deletion packages/core/src/config/config.ts
Original file line number Diff line number Diff line change
Expand Up @@ -3564,9 +3564,9 @@
// Gemini API key users should have the ability to manually select the
// old preview flash model.
if (authType === AuthType.USE_GEMINI) {
setFlashModels('gemini-3-flash-preview', 'gemini-3.5-flash');

Check warning on line 3567 in packages/core/src/config/config.ts

View workflow job for this annotation

GitHub Actions / Lint

Found sensitive keyword "gemini-3.5". Please make sure this change is appropriate to submit.
} else {
setFlashModels('gemini-3-flash', 'gemini-3-flash');
setFlashModels('gemini-3.5-flash', 'gemini-3.5-flash');

Check warning on line 3569 in packages/core/src/config/config.ts

View workflow job for this annotation

GitHub Actions / Lint

Found sensitive keyword "gemini-3.5". Please make sure this change is appropriate to submit.

Check warning on line 3569 in packages/core/src/config/config.ts

View workflow job for this annotation

GitHub Actions / Lint

Found sensitive keyword "gemini-3.5". Please make sure this change is appropriate to submit.
}
} else {
setFlashModels('gemini-3-flash-preview', 'gemini-2.5-flash');
Expand Down
4 changes: 4 additions & 0 deletions packages/core/src/config/models.ts
Original file line number Diff line number Diff line change
Expand Up @@ -59,14 +59,14 @@
// cleaned up.
export let PREVIEW_GEMINI_FLASH_MODEL = 'gemini-3-flash-preview';
export const DEFAULT_GEMINI_MODEL = 'gemini-2.5-pro';
// TODO: Set to const and update to 'gemini-3.5-flash' once the experiment for

Check warning on line 62 in packages/core/src/config/models.ts

View workflow job for this annotation

GitHub Actions / Lint

Found sensitive keyword "gemini-3.5". Please make sure this change is appropriate to submit.
// 3_5 flash rollut can be cleaned up.
// This is set to either the same as the DEFAULT_GEMINI_3_5_FLASH_MODEL const
// OR the SECONDARY_GEMINI_3_5_FLASH_MODEL depending on which is needed for
// the user's backend as determined by hasGemini35FlashGAAccess in
// packages/core/src/config/config.ts
export let DEFAULT_GEMINI_FLASH_MODEL = 'gemini-2.5-flash';
export const DEFAULT_GEMINI_3_5_FLASH_MODEL = 'gemini-3.5-flash';

Check warning on line 69 in packages/core/src/config/models.ts

View workflow job for this annotation

GitHub Actions / Lint

Found sensitive keyword "gemini-3.5". Please make sure this change is appropriate to submit.
// This is resolved to 3.5 flash in backends where it is used,
// however those backends do not expect to see the string gemini-3.5-flash
// so we need to provide this model as an alternative name in certain instances.
Expand Down Expand Up @@ -574,3 +574,7 @@
);
}
}

export const CCPA_AI_MODEL_MAPPINGS: Record<string, string> = {
[DEFAULT_GEMINI_3_5_FLASH_MODEL]: SECONDARY_GEMINI_3_5_FLASH_MODEL,
};
193 changes: 191 additions & 2 deletions packages/core/src/core/contentGenerator.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -18,10 +18,13 @@ import { HttpProxyAgent } from 'http-proxy-agent';
import { HttpsProxyAgent } from 'https-proxy-agent';
import type { Config } from '../config/config.js';
import { LoggingContentGenerator } from './loggingContentGenerator.js';
import { ModelMappingContentGenerator } from './modelMappingContentGenerator.js';
import { CCPA_AI_MODEL_MAPPINGS } from '../config/models.js';
import { loadApiKey } from './apiKeyCredentialStorage.js';
import { FakeContentGenerator } from './fakeContentGenerator.js';
import { RecordingContentGenerator } from './recordingContentGenerator.js';
import { resetVersionCache } from '../utils/version.js';
import type { LlmRole } from '../telemetry/llmRole.js';

vi.mock('../code_assist/codeAssist.js');
vi.mock('@google/genai');
Expand All @@ -36,6 +39,14 @@ const mockConfig = {
getProxy: vi.fn().mockReturnValue(undefined),
getUsageStatisticsEnabled: vi.fn().mockReturnValue(true),
getClientName: vi.fn().mockReturnValue(undefined),
getTelemetryLogPromptsEnabled: vi.fn().mockReturnValue(true),
getTelemetryTracesEnabled: vi.fn().mockReturnValue(true),
getSessionId: vi.fn().mockReturnValue('test-session-id'),
refreshUserQuotaIfStale: vi.fn().mockResolvedValue(undefined),
setLatestApiRequest: vi.fn(),
getContentGeneratorConfig: vi.fn().mockReturnValue({}),
isInteractive: vi.fn().mockReturnValue(false),
getExperiments: vi.fn().mockReturnValue(undefined),
} as unknown as Config;

describe('getAuthTypeFromEnv', () => {
Expand Down Expand Up @@ -142,7 +153,10 @@ describe('createContentGenerator', () => {
);
expect(createCodeAssistContentGenerator).toHaveBeenCalled();
expect(generator).toEqual(
new LoggingContentGenerator(mockGenerator, mockConfig),
new LoggingContentGenerator(
new ModelMappingContentGenerator(mockGenerator, CCPA_AI_MODEL_MAPPINGS),
mockConfig,
),
);
});

Expand All @@ -159,7 +173,10 @@ describe('createContentGenerator', () => {
);
expect(createCodeAssistContentGenerator).toHaveBeenCalled();
expect(generator).toEqual(
new LoggingContentGenerator(mockGenerator, mockConfig),
new LoggingContentGenerator(
new ModelMappingContentGenerator(mockGenerator, CCPA_AI_MODEL_MAPPINGS),
mockConfig,
),
);
});

Expand Down Expand Up @@ -1095,6 +1112,178 @@ describe('createContentGenerator', () => {
}),
);
});

it('should not apply model mapping for Vertex AI', async () => {
const mockModels = {
generateContent: vi.fn().mockResolvedValue({}),
};
const mockGenerator = {
models: mockModels,
} as unknown as GoogleGenAI;
vi.mocked(GoogleGenAI).mockImplementation(() => mockGenerator as never);

const generator = await createContentGenerator(
{
apiKey: 'test-api-key',
authType: AuthType.USE_VERTEX_AI,
vertexai: true,
},
mockConfig,
);

await generator.generateContent(
{
model: 'gemini-3-flash',
contents: [],
},
'prompt-id',
'user' as LlmRole,
);

expect(mockModels.generateContent).toHaveBeenCalledWith(
expect.objectContaining({
model: 'gemini-3-flash',
}),
'prompt-id',
'user',
);
});

it('should not apply model mapping for Gemini API', async () => {
const mockModels = {
generateContent: vi.fn().mockResolvedValue({}),
};
const mockGenerator = {
models: mockModels,
} as unknown as GoogleGenAI;
vi.mocked(GoogleGenAI).mockImplementation(() => mockGenerator as never);

const generator = await createContentGenerator(
{
apiKey: 'test-api-key',
authType: AuthType.USE_GEMINI,
},
mockConfig,
);

await generator.generateContent(
{
model: 'gemini-3-flash',
contents: [],
},
'prompt-id',
'user' as LlmRole,
);

expect(mockModels.generateContent).toHaveBeenCalledWith(
expect.objectContaining({
model: 'gemini-3-flash',
}),
'prompt-id',
'user',
);
});

it('should not apply model mapping for GATEWAY', async () => {
const mockModels = {
generateContent: vi.fn().mockResolvedValue({}),
};
const mockGenerator = {
models: mockModels,
} as unknown as GoogleGenAI;
vi.mocked(GoogleGenAI).mockImplementation(() => mockGenerator as never);

const generator = await createContentGenerator(
{
apiKey: 'test-api-key',
authType: AuthType.GATEWAY,
},
mockConfig,
);

await generator.generateContent(
{
model: 'gemini-3.5-flash',
contents: [],
},
'prompt-id',
'user' as LlmRole,
);

expect(mockModels.generateContent).toHaveBeenCalledWith(
expect.objectContaining({
model: 'gemini-3.5-flash',
}),
'prompt-id',
'user',
);
});

it('should apply model mapping for LOGIN_WITH_GOOGLE', async () => {
const mockInnerGenerator = {
generateContent: vi.fn().mockResolvedValue({}),
} as unknown as ContentGenerator;
vi.mocked(createCodeAssistContentGenerator).mockResolvedValue(
mockInnerGenerator as never,
);

const generator = await createContentGenerator(
{
authType: AuthType.LOGIN_WITH_GOOGLE,
},
mockConfig,
);

await generator.generateContent(
{
model: 'gemini-3.5-flash',
contents: [],
},
'prompt-id',
'user' as LlmRole,
);

expect(mockInnerGenerator.generateContent).toHaveBeenCalledWith(
expect.objectContaining({
model: 'gemini-3-flash',
}),
'prompt-id',
'user',
);
});

it('should apply model mapping for COMPUTE_ADC', async () => {
const mockInnerGenerator = {
generateContent: vi.fn().mockResolvedValue({}),
} as unknown as ContentGenerator;
vi.mocked(createCodeAssistContentGenerator).mockResolvedValue(
mockInnerGenerator as never,
);

const generator = await createContentGenerator(
{
authType: AuthType.COMPUTE_ADC,
},
mockConfig,
);

await generator.generateContent(
{
model: 'gemini-3.5-flash',
contents: [],
},
'prompt-id',
'user' as LlmRole,
);

expect(mockInnerGenerator.generateContent).toHaveBeenCalledWith(
expect.objectContaining({
model: 'gemini-3-flash',
}),
'prompt-id',
'user',
);
});
});

describe('createContentGeneratorConfig', () => {
Expand Down
15 changes: 10 additions & 5 deletions packages/core/src/core/contentGenerator.ts
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,8 @@ import { determineSurface } from '../utils/surface.js';
import { RecordingContentGenerator } from './recordingContentGenerator.js';
import { getVersion, resolveModel } from '../../index.js';
import type { LlmRole } from '../telemetry/llmRole.js';
import { ModelMappingContentGenerator } from './modelMappingContentGenerator.js';
import { CCPA_AI_MODEL_MAPPINGS } from '../config/models.js';

/**
* Interface abstracting the core functionalities for generating content and counting tokens.
Expand Down Expand Up @@ -282,11 +284,14 @@ export async function createContentGenerator(
) {
const httpOptions = { headers: baseHeaders };
return new LoggingContentGenerator(
await createCodeAssistContentGenerator(
httpOptions,
config.authType,
gcConfig,
sessionId,
new ModelMappingContentGenerator(
await createCodeAssistContentGenerator(
httpOptions,
config.authType,
gcConfig,
sessionId,
),
CCPA_AI_MODEL_MAPPINGS,
),
gcConfig,
);
Expand Down
Loading
Loading