Skip to content
64 changes: 64 additions & 0 deletions packages/core/src/followup/speculation.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -101,6 +101,70 @@ describe('startSpeculation', () => {
});
});

describe.each([
{
scenario: 'same model (undefined)',
fastModel: undefined,
expectedPreserveTools: true,
},
{
scenario: 'different model',
fastModel: 'different-fast-model',
expectedPreserveTools: false,
},
])(
'generatePipelinedSuggestion preserveTools — $scenario',
({ fastModel, expectedPreserveTools }) => {
it(`passes preserveTools: ${String(expectedPreserveTools)} to runForkedAgent`, async () => {
const config = {
getApprovalMode: vi.fn().mockReturnValue(ApprovalMode.DEFAULT),
getCwd: vi.fn().mockReturnValue(process.cwd()),
getFastModel: vi.fn().mockReturnValue(fastModel),
getToolRegistry: vi.fn().mockReturnValue({
ensureTool: vi.fn().mockResolvedValue({
build: vi.fn().mockReturnValue({
execute: vi.fn().mockResolvedValue({
llmContent: '',
returnDisplay: '',
}),
}),
}),
}),
} as unknown as Config;

forkedAgentMocks.runForkedAgent.mockResolvedValue({
jsonResult: { suggestion: 'next step' },
});

forkedAgentMocks.sendMessageStream.mockImplementation(async function* () {
yield {
type: 'chunk',
value: {
candidates: [
{
content: {
parts: [{ text: 'done' }],
},
},
],
},
};
});

const state = await startSpeculation(config, 'do something');
await vi.waitFor(() => {
expect(state.status).toBe('completed');
});

expect(forkedAgentMocks.runForkedAgent).toHaveBeenCalledWith(
expect.objectContaining({ preserveTools: expectedPreserveTools }),
);

await abortSpeculation(state);
});
},
);

describe('ensureToolResultPairing', () => {
it('returns empty array unchanged', () => {
expect(ensureToolResultPairing([])).toEqual([]);
Expand Down
2 changes: 2 additions & 0 deletions packages/core/src/followup/speculation.ts
Original file line number Diff line number Diff line change
Expand Up @@ -554,13 +554,15 @@ ${SUGGESTION_PROMPT}`;
const cacheSafeParams = getCacheSafeParams();
if (!cacheSafeParams) return null;
const model = modelOverride ?? config.getFastModel();
const resolvedModel = model ?? cacheSafeParams.model;
const result = await runForkedAgent({
config,
userMessage: augmentedPrompt,
cacheSafeParams,
jsonSchema: PIPELINED_SCHEMA,
...(model !== undefined ? { model } : {}),
abortSignal,
preserveTools: resolvedModel === cacheSafeParams.model,
});

if (abortSignal.aborted) return null;
Expand Down
57 changes: 57 additions & 0 deletions packages/core/src/followup/suggestionGenerator.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,63 @@ describe('generatePromptSuggestion', () => {
expect.objectContaining({ model: 'openai:fast-model' }),
);
});
it('passes preserveTools: true for Anthropic prompt-cache sharing', async () => {
mockGetCacheSafeParams.mockReturnValue({
generationConfig: {},
history: conversationHistory,
model: 'main-model',
version: 1,
});
mockRunForkedAgent.mockResolvedValue({
text: null,
jsonResult: { suggestion: 'run tests' },
usage: { inputTokens: 10, outputTokens: 3, cacheHitTokens: 5 },
});
const config = {
getFastModel: vi.fn(() => undefined),
getModel: vi.fn(() => 'main-model'),
} as unknown as Config;

await generatePromptSuggestion(
config,
conversationHistory,
new AbortController().signal,
{ enableCacheSharing: true },
);

expect(mockRunForkedAgent).toHaveBeenCalledWith(
expect.objectContaining({ preserveTools: true }),
);
});

it('passes preserveTools: false when fast model differs from cache-safe model', async () => {
mockGetCacheSafeParams.mockReturnValue({
generationConfig: {},
history: conversationHistory,
model: 'main-model',
version: 1,
});
mockRunForkedAgent.mockResolvedValue({
text: null,
jsonResult: { suggestion: 'run tests' },
usage: { inputTokens: 10, outputTokens: 3, cacheHitTokens: 5 },
});
const config = {
getFastModel: vi.fn(() => 'different-fast-model'),
getModel: vi.fn(() => 'main-model'),
} as unknown as Config;

await generatePromptSuggestion(
config,
conversationHistory,
new AbortController().signal,
{ enableCacheSharing: true },
);

expect(mockRunForkedAgent).toHaveBeenCalledWith(
expect.objectContaining({ preserveTools: false }),
);
});
});

describe('shouldFilterSuggestion', () => {
Expand Down
1 change: 1 addition & 0 deletions packages/core/src/followup/suggestionGenerator.ts
Original file line number Diff line number Diff line change
Expand Up @@ -156,6 +156,7 @@ async function generateViaForkedQuery(
cacheSafeParams,
jsonSchema: SUGGESTION_SCHEMA,
model,
preserveTools: model === cacheSafeParams.model,
});

if (result.jsonResult) {
Expand Down
Loading
Loading