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/cli/src/ui/components/RewindViewer.test.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -426,6 +426,53 @@ describe('RewindViewer', () => {
expect(lastFrame2()).toMatchSnapshot('after-update');
unmount2();
});

it('excludes user messages that only carry tool responses from the rewind points', async () => {
const messages: MessageRecord[] = [
{ type: 'user', content: 'Run the tool', id: '1', timestamp: '1' },
{ type: 'gemini', content: 'Running it now.', id: '2', timestamp: '2' },
{
type: 'user',
content: [
{
functionResponse: {
name: 'testTool',
id: 'call-1',
response: { output: 'tool output' },
},
},
],
id: '3',
timestamp: '3',
},
{
type: 'user',
content: 'Thanks, next question',
id: '4',
timestamp: '4',
},
];
const conversation = createConversation(messages);
const onExit = vi.fn();
const onRewind = vi.fn();

const { lastFrame, unmount } = await renderWithProviders(
<RewindViewer
conversation={conversation}
onExit={onExit}
onRewind={onRewind}
/>,
);

const frame = lastFrame();
expect(frame).toContain('Run the tool');
expect(frame).toContain('Thanks, next question');
expect(frame).toContain('Stay at current position');
// Only the two authored prompts are listed as rewind points; the
// tool-response record between them is not offered as a third entry.
expect((frame?.match(/No files have been changed/g) ?? []).length).toBe(2);
unmount();
});
});
it('renders accessible screen reader view when screen reader is enabled', async () => {
const { useIsScreenReaderEnabled } = await import('ink');
Expand Down
6 changes: 5 additions & 1 deletion packages/cli/src/ui/components/RewindViewer.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@ import { useKeypress } from '../hooks/useKeypress.js';
import { useRewind } from '../hooks/useRewind.js';
import { RewindConfirmation, RewindOutcome } from './RewindConfirmation.js';
import { stripReferenceContent } from '../utils/formatters.js';
import { isToolResponseMessage } from '../utils/rewindFileOps.js';
import { Command } from '../key/keyMatchers.js';
import { CliSpinner } from './CliSpinner.js';
import { ExpandableText } from './shared/ExpandableText.js';
Expand Down Expand Up @@ -69,7 +70,10 @@ export const RewindViewer: React.FC<RewindViewerProps> = ({
);

const interactions = useMemo(
() => conversation.messages.filter((msg) => msg.type === 'user'),
() =>
conversation.messages.filter(
(msg) => msg.type === 'user' && !isToolResponseMessage(msg),
),
[conversation.messages],
);
Comment thread
luisfelipe-alt marked this conversation as resolved.

Expand Down
172 changes: 170 additions & 2 deletions packages/cli/src/ui/hooks/useGeminiStream.test.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,19 @@ const MockedGeminiClientClass = vi.hoisted(() =>
this.setHistory = vi.fn().mockImplementation((newHistory: any[]) => {
mockHistory = [...newHistory];
});
this.discardTrailingUnansweredToolCallTurn = vi
.fn()
.mockImplementation(() => {
const last = mockHistory[mockHistory.length - 1];
if (
last?.role !== 'model' ||
!last.parts?.some((part: any) => !!part.functionCall)
) {
return false;
}
mockHistory = mockHistory.slice(0, -1);
return true;
});
this.generateContent = vi.fn().mockResolvedValue({
candidates: [
{ content: { parts: [{ text: 'Got it. Focusing on tests only.' }] } },
Expand Down Expand Up @@ -1132,7 +1145,7 @@ describe('useGeminiStream', () => {
),
);

// Call submitQuery to populate the user turn and set historyLengthAfterUserPromptRef
// Call submitQuery to populate the user turn
await act(async () => {
// eslint-disable-next-line @typescript-eslint/no-floating-promises
result.current.submitQuery('User prompt');
Expand Down Expand Up @@ -1167,6 +1180,161 @@ describe('useGeminiStream', () => {
});
});

it('should keep the user prompt and completed tool rounds when a later tool batch is declined', async () => {
const cancelledToolCalls: TrackedToolCall[] = [
{
request: {
callId: '2',
name: 'testTool',
args: {},
isClientInitiated: false,
prompt_id: 'prompt-id-4',
},
status: CoreToolCallStatus.Cancelled,
response: {
callId: '2',
responseParts: [{ text: CoreToolCallStatus.Cancelled }],
errorType: undefined,
},
responseSubmittedToGemini: false,
tool: {
displayName: 'mock tool',
},
invocation: {
getDescription: () => `Mock description`,
},
} as any,
];
const client = new MockedGeminiClientClass(mockConfig);
const priorTurn = [
{ role: 'user', parts: [{ text: 'Earlier prompt' }] },
{ role: 'model', parts: [{ text: 'Earlier answer' }] },
];
client.setHistory(priorTurn);
// Model the real sendMessageStream contract: the user turn is recorded in
// the client history when the request is sent, not before submitQuery.
mockSendMessageStream.mockImplementation((query: PartListUnion) => {
const parts = (Array.isArray(query) ? query : [query]).map((part) =>
typeof part === 'string' ? { text: part } : part,
);
client.setHistory([...client.getHistory(), { role: 'user', parts }]);
return (async function* () {
yield { type: ServerGeminiEventType.Content, value: 'Working on it' };
})();
});

let capturedOnComplete:
| ((completedTools: TrackedToolCall[]) => Promise<void>)
| null = null;

mockUseToolScheduler.mockImplementation((onComplete) => {
capturedOnComplete = onComplete;
return [
[],
mockScheduleToolCalls,
mockMarkToolsAsSubmitted,
vi.fn(),
mockCancelAllToolCalls,
0,
];
});

const { result } = await renderHookWithProviders(() =>
useGeminiStream(
client,
[],
mockAddItem,
mockConfig,
mockLoadedSettings,
mockOnDebugMessage,
mockHandleSlashCommand,
false,
() => 'vscode' as EditorType,
() => {},
() => Promise.resolve(),
false,
() => {},
() => {},
() => {},
80,
24,
),
);

// Turn 1 starts: sendMessageStream records the user prompt.
await act(async () => {
await result.current.submitQuery('User prompt');
});

// The model answers with a first tool call, which completes and is sent
// back to the model as a continuation turn.
const firstRoundResponse: Part[] = [
{
functionResponse: {
name: 'testTool',
id: '1',
response: { output: 'first result' },
},
},
];
client.setHistory([
...client.getHistory(),
{
role: 'model',
parts: [{ functionCall: { name: 'testTool', id: '1', args: {} } }],
},
]);
await act(async () => {
await result.current.submitQuery(firstRoundResponse, {
isContinuation: true,
});
});

// The model requests a second tool call within the same chain.
client.setHistory([
...client.getHistory(),
{
role: 'model',
parts: [{ functionCall: { name: 'testTool', id: '2', args: {} } }],
},
]);
vi.mocked(client.setHistory).mockClear();

// The second tool call is declined.
await act(async () => {
if (capturedOnComplete) {
await new Promise((resolve) => setTimeout(resolve, 0));
await capturedOnComplete(cancelledToolCalls);
}
});

try {
await waitFor(() => {
expect(mockMarkToolsAsSubmitted).toHaveBeenCalledWith(['2']);
expect(client.addHistory).not.toHaveBeenCalled();
// Only the trailing unanswered call is removed, without re-setting
// the whole history (which would re-record every turn).
expect(
client.discardTrailingUnansweredToolCallTurn,
).toHaveBeenCalledTimes(1);
expect(client.setHistory).not.toHaveBeenCalled();
// The previous turn, the user prompt and the completed first round
// are preserved.
expect(client.getHistory()).toEqual([
...priorTurn,
{ role: 'user', parts: [{ text: 'User prompt' }] },
{
role: 'model',
parts: [{ functionCall: { name: 'testTool', id: '1', args: {} } }],
},
{ role: 'user', parts: firstRoundResponse },
]);
});
} finally {
mockSendMessageStream.mockImplementation(() => (async function* () {})());
}
});

it('should record tool responses in history when the model was switched due to a quota error', async () => {
// Regression test: returning early on a quota-triggered model switch
// without recording the responses leaves the already-recorded
Expand Down Expand Up @@ -1597,7 +1765,7 @@ describe('useGeminiStream', () => {
),
);

// Call submitQuery to populate the user turn and set historyLengthAfterUserPromptRef
// Call submitQuery to populate the user turn
await act(async () => {
// eslint-disable-next-line @typescript-eslint/no-floating-promises
result.current.submitQuery('User prompt');
Expand Down
21 changes: 4 additions & 17 deletions packages/cli/src/ui/hooks/useGeminiStream.ts
Original file line number Diff line number Diff line change
Expand Up @@ -258,7 +258,6 @@ export const useGeminiStream = (
const abortControllerRef = useRef<AbortController | null>(null);
const turnCancelledRef = useRef(false);
const activeQueryIdRef = useRef<string | null>(null);
const historyLengthAfterUserPromptRef = useRef<number | undefined>(undefined);
const previousApprovalModeRef = useRef<ApprovalMode>(
config.getApprovalMode(),
);
Expand Down Expand Up @@ -1739,11 +1738,6 @@ export const useGeminiStream = (
return;
}

if (geminiClient) {
historyLengthAfterUserPromptRef.current =
geminiClient.getHistory().length;
}

if (!options?.isContinuation) {
if (typeof queryToSend === 'string') {
// logging the text prompts only for now
Expand Down Expand Up @@ -2127,17 +2121,10 @@ export const useGeminiStream = (
}
setIsResponding(false);

if (
geminiClient &&
historyLengthAfterUserPromptRef.current !== undefined
) {
const targetLength = historyLengthAfterUserPromptRef.current;
if (geminiClient.getHistory().length > targetLength) {
geminiClient.setHistory(
geminiClient.getHistory().slice(0, targetLength),
);
}
}
// Only roll back the unanswered model function call turn for this
// cancelled batch. The originating user prompt and any tool rounds
// that already completed (and may have changed files) stay in history.
geminiClient?.discardTrailingUnansweredToolCallTurn();

const callIdsToMarkAsSubmitted = geminiTools.map(
(toolCall) => toolCall.request.callId,
Expand Down
Loading
Loading