diff --git a/packages/ai-ollama/src/node/ollama-language-model.spec.ts b/packages/ai-ollama/src/node/ollama-language-model.spec.ts index cf4175d2f36b9..f38a190a14c76 100644 --- a/packages/ai-ollama/src/node/ollama-language-model.spec.ts +++ b/packages/ai-ollama/src/node/ollama-language-model.spec.ts @@ -15,7 +15,7 @@ // ***************************************************************************** import { expect } from 'chai'; -import { Message } from 'ollama'; +import { Message, Ollama } from 'ollama'; import { OllamaModel } from './ollama-language-model'; class TestableOllamaModel extends OllamaModel { @@ -26,6 +26,15 @@ class TestableOllamaModel extends OllamaModel { public callMergeConsecutiveAssistantMessages(messages: Message[]): Message[] { return this.mergeConsecutiveAssistantMessages(messages); } + + public callHandleStreamingRequest(ollama: Ollama, messages: Message[]): ReturnType { + return this.handleStreamingRequest(ollama, { + model: 'test-model', + messages, + stream: true + }); + } + } describe('OllamaModel - mergeConsecutiveAssistantMessages', () => { @@ -121,3 +130,49 @@ describe('OllamaModel - mergeConsecutiveAssistantMessages', () => { expect(result).to.deep.equal(messages); }); }); + +describe('OllamaModel - handleStreamingRequest', () => { + + it('should accept tool_calls as a valid done reason', async () => { + const model = new TestableOllamaModel(); + + const responseStream = { + async *[Symbol.asyncIterator](): AsyncGenerator<{ + created_at: Date; + done: boolean; + done_reason: string; + message: { role: string; content: string }; + }> { + yield { + created_at: new Date(), + done: true, + done_reason: 'tool_calls', + message: { + role: 'assistant', + content: '' + } + }; + }, + abort: () => undefined + }; + + const ollama = { + show: async () => ({ capabilities: [] }), + chat: async () => responseStream + } as unknown as Ollama; + + const response = await model.callHandleStreamingRequest(ollama, [ + { role: 'user', content: 'use a tool' } + ]); + + const parts = []; + + for await (const part of response.stream) { + parts.push(part); + } + + expect(parts).to.deep.equal([]); + }); + +}); + diff --git a/packages/ai-ollama/src/node/ollama-language-model.ts b/packages/ai-ollama/src/node/ollama-language-model.ts index 7d81442ebac23..fde09dadaa55e 100644 --- a/packages/ai-ollama/src/node/ollama-language-model.ts +++ b/packages/ai-ollama/src/node/ollama-language-model.ts @@ -172,7 +172,7 @@ export class OllamaModel implements LanguageModel { if (chunk.prompt_eval_count !== undefined && chunk.eval_count !== undefined) { yield { input_tokens: chunk.prompt_eval_count, output_tokens: chunk.eval_count }; } - if (chunk.done_reason && chunk.done_reason !== 'stop') { + if (chunk.done_reason && chunk.done_reason !== 'stop' && chunk.done_reason !== 'tool_calls') { throw new Error('Ollama stopped unexpectedly. Reason: ' + chunk.done_reason); } } @@ -311,7 +311,7 @@ export class OllamaModel implements LanguageModel { lastUpdated = chunk.created_at; inputTokenCount = chunk.prompt_eval_count; outputTokenCount = chunk.eval_count; - if (chunk.done_reason && chunk.done_reason !== 'stop') { + if (chunk.done_reason && chunk.done_reason !== 'stop' && chunk.done_reason !== 'tool_calls') { throw new Error('Ollama stopped unexpectedly. Reason: ' + chunk.done_reason); } }