1use crate::error::Result;
4use crate::{
5 Error,
6 core::{
7 AssistantMessage, Message,
8 language_model::{
9 LanguageModel, LanguageModelOptions, LanguageModelResponse,
10 LanguageModelResponseContentType, StopReason, request::LanguageModelRequest,
11 },
12 messages::TaggedMessage,
13 utils::resolve_message,
14 },
15};
16use serde::de::DeserializeOwned;
17use serde::ser::Error as SerdeError;
18use std::ops::Deref;
19
20impl<M: LanguageModel> LanguageModelRequest<M> {
21 pub async fn generate_text(&mut self) -> Result<GenerateTextResponse> {
67 let (system_prompt, messages) = resolve_message(&self.options, &self.prompt);
68
69 let mut options = LanguageModelOptions {
70 system: (!system_prompt.is_empty()).then_some(system_prompt),
71 messages,
72 schema: self.options.schema.to_owned(),
73 stop_sequences: self.options.stop_sequences.to_owned(),
74 tools: self.options.tools.to_owned(),
75 stop_when: self.options.stop_when.clone(),
76 on_step_start: self.options.on_step_start.clone(),
77 on_step_finish: self.options.on_step_finish.clone(),
78 stop_reason: None,
79 ..self.options
80 };
81
82 loop {
83 options.current_step_id += 1;
85
86 if let Some(hook) = options.on_step_start.clone() {
88 hook(&mut options);
89 }
90
91 let response: LanguageModelResponse = self
92 .model
93 .generate_text(options.clone())
94 .await
95 .inspect_err(|e| {
96 options.stop_reason = Some(StopReason::Error(e.clone()));
97 })?;
98
99 for output in response.contents.iter() {
100 match output {
101 LanguageModelResponseContentType::Text(text) => {
102 let assistant_msg = Message::Assistant(AssistantMessage {
103 content: text.clone().into(),
104 usage: response.usage.clone(),
105 });
106 options
107 .messages
108 .push(TaggedMessage::new(options.current_step_id, assistant_msg));
109 }
110 LanguageModelResponseContentType::Reasoning {
111 content,
112 extensions,
113 } => {
114 let assistant_msg = Message::Assistant(AssistantMessage {
115 content: LanguageModelResponseContentType::Reasoning {
116 content: content.clone(),
117 extensions: extensions.clone(),
118 },
119 usage: response.usage.clone(),
120 });
121 options
122 .messages
123 .push(TaggedMessage::new(options.current_step_id, assistant_msg));
124 }
125 LanguageModelResponseContentType::ToolCall(tool_info) => {
126 let usage = response.usage.clone();
128 let _ = &options.messages.push(TaggedMessage::new(
129 options.current_step_id.to_owned(),
130 Message::Assistant(AssistantMessage::new(
131 LanguageModelResponseContentType::ToolCall(tool_info.clone()),
132 usage,
133 )),
134 ));
135 options.handle_tool_call(tool_info).await;
136 }
137 _ => (),
138 }
139 }
140
141 if let Some(ref hook) = options.on_step_finish {
143 hook(&options);
144 };
145
146 if response.contents.is_empty() {
147 options.stop_reason = Some(StopReason::Error(Error::Other(
148 "Language model returned empty response".to_string(),
149 )));
150 break;
151 }
152
153 if let Some(hook) = &options.stop_when.clone()
155 && hook(&options)
156 {
157 options.stop_reason = Some(StopReason::Hook);
158 break;
159 }
160
161 match response.contents.last() {
162 Some(LanguageModelResponseContentType::ToolCall(_)) => (),
163 _ => {
164 options.stop_reason = Some(StopReason::Finish);
165 break;
166 }
167 };
168 }
169
170 Ok(GenerateTextResponse { options })
171 }
172}
173
174#[derive(Debug, Clone)]
180pub struct GenerateTextResponse {
181 pub options: LanguageModelOptions,
183}
184
185impl GenerateTextResponse {
186 pub fn into_schema<T: DeserializeOwned>(&self) -> std::result::Result<T, serde_json::Error> {
203 if let Some(text) = &self.text() {
204 serde_json::from_str(text)
205 } else {
206 Err(serde_json::Error::custom("No text response found"))
207 }
208 }
209
210 #[cfg(any(test, feature = "test-access"))]
211 pub fn step_ids(&self) -> Vec<usize> {
213 self.options.messages.iter().map(|t| t.step_id).collect()
214 }
215}
216
217impl Deref for GenerateTextResponse {
218 type Target = LanguageModelOptions;
219
220 fn deref(&self) -> &Self::Target {
221 &self.options
222 }
223}
224
225#[cfg(test)]
226mod tests {
227 use super::*;
228 use crate::core::{
229 AssistantMessage,
230 language_model::{LanguageModelResponseContentType, Usage},
231 messages::TaggedMessage,
232 tools::{ToolCallInfo, ToolResultInfo},
233 };
234
235 #[test]
236 fn test_generate_text_response_step() {
237 let options = LanguageModelOptions {
238 messages: vec![
239 TaggedMessage::new(0, Message::System("System".to_string().into())),
240 TaggedMessage::new(0, Message::User("User".to_string().into())),
241 TaggedMessage::new(
242 1,
243 Message::Assistant(AssistantMessage {
244 content: LanguageModelResponseContentType::Text("Assistant".to_string()),
245 usage: None,
246 }),
247 ),
248 ],
249 ..Default::default()
250 };
251 let response = GenerateTextResponse { options };
252
253 let step0 = response.step(0).unwrap();
254 assert_eq!(step0.step_id, 0);
255 assert_eq!(step0.messages.len(), 2);
256
257 let step1 = response.step(1).unwrap();
258 assert_eq!(step1.step_id, 1);
259 assert_eq!(step1.messages.len(), 1);
260
261 assert!(response.step(2).is_none());
262 }
263
264 #[test]
265 fn test_generate_text_response_final_step() {
266 let options = LanguageModelOptions {
267 messages: vec![
268 TaggedMessage::new(0, Message::System("System".to_string().into())),
269 TaggedMessage::new(1, Message::User("User".to_string().into())),
270 TaggedMessage::new(
271 2,
272 Message::Assistant(AssistantMessage {
273 content: LanguageModelResponseContentType::Text("Assistant".to_string()),
274 usage: None,
275 }),
276 ),
277 ],
278 ..Default::default()
279 };
280 let response = GenerateTextResponse { options };
281
282 let final_step = response.last_step().unwrap();
283 assert_eq!(final_step.step_id, 2);
284 assert_eq!(final_step.messages.len(), 1);
285 }
286
287 #[test]
288 fn test_generate_text_response_steps() {
289 let options = LanguageModelOptions {
290 messages: vec![
291 TaggedMessage::new(0, Message::System("System".to_string().into())),
292 TaggedMessage::new(0, Message::User("User".to_string().into())),
293 TaggedMessage::new(
294 1,
295 Message::Assistant(AssistantMessage {
296 content: LanguageModelResponseContentType::Text("Assistant1".to_string()),
297 usage: None,
298 }),
299 ),
300 TaggedMessage::new(
301 2,
302 Message::Assistant(AssistantMessage {
303 content: LanguageModelResponseContentType::Text("Assistant2".to_string()),
304 usage: None,
305 }),
306 ),
307 ],
308 ..Default::default()
309 };
310 let response = GenerateTextResponse { options };
311
312 let steps = response.steps();
313 assert_eq!(steps.len(), 3);
314 assert_eq!(steps[0].step_id, 0);
315 assert_eq!(steps[0].messages.len(), 2);
316 assert_eq!(steps[1].step_id, 1);
317 assert_eq!(steps[1].messages.len(), 1);
318 assert_eq!(steps[2].step_id, 2);
319 assert_eq!(steps[2].messages.len(), 1);
320 }
321
322 #[test]
323 fn test_generate_text_response_usage() {
324 let options = LanguageModelOptions {
325 messages: vec![
326 TaggedMessage::new(0, Message::System("System".to_string().into())),
327 TaggedMessage::new(
328 1,
329 Message::Assistant(AssistantMessage {
330 content: LanguageModelResponseContentType::Text("Assistant1".to_string()),
331 usage: Some(Usage {
332 input_tokens: Some(10),
333 output_tokens: Some(5),
334 reasoning_tokens: Some(2),
335 cached_tokens: Some(1),
336 }),
337 }),
338 ),
339 TaggedMessage::new(
340 2,
341 Message::Assistant(AssistantMessage {
342 content: LanguageModelResponseContentType::Text("Assistant2".to_string()),
343 usage: Some(Usage {
344 input_tokens: Some(5),
345 output_tokens: Some(3),
346 reasoning_tokens: Some(1),
347 cached_tokens: Some(0),
348 }),
349 }),
350 ),
351 ],
352 ..Default::default()
353 };
354 let response = GenerateTextResponse { options };
355
356 let total_usage = response.usage();
357 assert_eq!(total_usage.input_tokens, Some(15));
358 assert_eq!(total_usage.output_tokens, Some(8));
359 assert_eq!(total_usage.reasoning_tokens, Some(3));
360 assert_eq!(total_usage.cached_tokens, Some(1));
361 }
362
363 fn create_tool_call_message(step_id: usize, tool_name: &str) -> TaggedMessage {
364 TaggedMessage::new(
365 step_id,
366 Message::Assistant(AssistantMessage {
367 content: LanguageModelResponseContentType::ToolCall(ToolCallInfo::new(tool_name)),
368 usage: None,
369 }),
370 )
371 }
372
373 fn create_tool_result_message(step_id: usize, tool_name: &str) -> TaggedMessage {
374 TaggedMessage::new(step_id, Message::Tool(ToolResultInfo::new(tool_name)))
375 }
376
377 fn create_text_assistant_message(step_id: usize, text: &str) -> TaggedMessage {
378 TaggedMessage::new(
379 step_id,
380 Message::Assistant(AssistantMessage {
381 content: LanguageModelResponseContentType::Text(text.to_string()),
382 usage: None,
383 }),
384 )
385 }
386
387 fn create_response_with_messages(messages: Vec<TaggedMessage>) -> GenerateTextResponse {
388 let options = LanguageModelOptions {
389 messages,
390 ..Default::default()
391 };
392 GenerateTextResponse { options }
393 }
394
395 #[test]
397 fn test_generate_text_response_tool_calls_empty_messages() {
398 let response = create_response_with_messages(vec![]);
399 assert_eq!(response.tool_calls(), None);
400 }
401
402 #[test]
403 fn test_generate_text_response_tool_calls_only_non_assistant_messages() {
404 let messages = vec![
405 TaggedMessage::new(0, Message::System("System".to_string().into())),
406 TaggedMessage::new(0, Message::User("User".to_string().into())),
407 create_tool_result_message(0, "tool1"),
408 ];
409 let response = create_response_with_messages(messages);
410 assert_eq!(response.tool_calls(), None);
411 }
412
413 #[test]
414 fn test_generate_text_response_tool_calls_single_assistant_with_tool_call() {
415 let messages = vec![create_tool_call_message(0, "test_tool")];
416 let response = create_response_with_messages(messages);
417 let calls = response.tool_calls().unwrap();
418 assert_eq!(calls.len(), 1);
419 assert_eq!(calls[0].tool.name, "test_tool");
420 }
421
422 #[test]
423 fn test_generate_text_response_tool_calls_multiple_assistant_with_tool_calls_different_steps() {
424 let messages = vec![
425 create_tool_call_message(0, "tool1"),
426 create_tool_call_message(1, "tool2"),
427 create_tool_call_message(2, "tool3"),
428 ];
429 let response = create_response_with_messages(messages);
430 let calls = response.tool_calls().unwrap();
431 assert_eq!(calls.len(), 3);
432 assert_eq!(calls[0].tool.name, "tool1");
433 assert_eq!(calls[1].tool.name, "tool2");
434 assert_eq!(calls[2].tool.name, "tool3");
435 }
436
437 #[test]
438 fn test_generate_text_response_tool_calls_assistant_without_tool_call() {
439 let messages = vec![create_text_assistant_message(0, "Hello")];
440 let response = create_response_with_messages(messages);
441 assert_eq!(response.tool_calls(), None);
442 }
443
444 #[test]
445 fn test_generate_text_response_tool_calls_mixed_message_types_multiple_steps() {
446 let messages = vec![
447 TaggedMessage::new(0, Message::System("System".to_string().into())),
448 TaggedMessage::new(0, Message::User("User".to_string().into())),
449 create_tool_call_message(1, "test_tool"),
450 create_tool_result_message(1, "other_tool"),
451 create_tool_call_message(2, "another_tool"),
452 ];
453 let response = create_response_with_messages(messages);
454 let calls = response.tool_calls().unwrap();
455 assert_eq!(calls.len(), 2);
456 assert_eq!(calls[0].tool.name, "test_tool");
457 assert_eq!(calls[1].tool.name, "another_tool");
458 }
459
460 #[test]
461 fn test_generate_text_response_tool_calls_duplicate_tool_calls() {
462 let messages = vec![
463 create_tool_call_message(0, "tool1"),
464 create_tool_call_message(1, "tool1"), create_tool_call_message(2, "tool1"), ];
467 let response = create_response_with_messages(messages);
468 let calls = response.tool_calls().unwrap();
469 assert_eq!(calls.len(), 3);
470 assert_eq!(calls[0].tool.name, "tool1");
471 assert_eq!(calls[1].tool.name, "tool1");
472 assert_eq!(calls[2].tool.name, "tool1");
473 }
474
475 #[test]
476 fn test_generate_text_response_tool_calls_from_specific_steps_only() {
477 let messages = vec![
478 TaggedMessage::new(0, Message::System("System".to_string().into())),
479 create_tool_call_message(1, "tool_from_step1"),
480 TaggedMessage::new(1, Message::User("User".to_string().into())),
481 create_tool_call_message(2, "tool_from_step2"),
482 create_tool_result_message(2, "result_from_step2"),
483 create_tool_call_message(3, "tool_from_step3"),
484 ];
485 let response = create_response_with_messages(messages);
486 let calls = response.tool_calls().unwrap();
487 assert_eq!(calls.len(), 3);
488 assert_eq!(calls[0].tool.name, "tool_from_step1");
489 assert_eq!(calls[1].tool.name, "tool_from_step2");
490 assert_eq!(calls[2].tool.name, "tool_from_step3");
491 }
492
493 #[test]
495 fn test_generate_text_response_tool_results_empty_messages() {
496 let response = create_response_with_messages(vec![]);
497 assert!(response.tool_results().is_none());
498 }
499
500 #[test]
501 fn test_generate_text_response_tool_results_only_non_tool_messages() {
502 let messages = vec![
503 TaggedMessage::new(0, Message::System("System".to_string().into())),
504 TaggedMessage::new(0, Message::User("User".to_string().into())),
505 create_text_assistant_message(0, "Assistant"),
506 ];
507 let response = create_response_with_messages(messages);
508 assert!(response.tool_results().is_none());
509 }
510
511 #[test]
512 fn test_generate_text_response_tool_results_single_tool_message() {
513 let messages = vec![create_tool_result_message(0, "test_tool")];
514 let response = create_response_with_messages(messages);
515 let results = response.tool_results().unwrap();
516 assert_eq!(results.len(), 1);
517 assert_eq!(results[0].tool.name, "test_tool");
518 }
519
520 #[test]
521 fn test_generate_text_response_tool_results_multiple_tool_messages_different_steps() {
522 let messages = vec![
523 create_tool_result_message(0, "tool1"),
524 create_tool_result_message(1, "tool2"),
525 create_tool_result_message(2, "tool3"),
526 ];
527 let response = create_response_with_messages(messages);
528 let results = response.tool_results().unwrap();
529 assert_eq!(results.len(), 3);
530 assert_eq!(results[0].tool.name, "tool1");
531 assert_eq!(results[1].tool.name, "tool2");
532 assert_eq!(results[2].tool.name, "tool3");
533 }
534
535 #[test]
536 fn test_generate_text_response_tool_results_mixed_message_types() {
537 let messages = vec![
538 TaggedMessage::new(0, Message::System("System".to_string().into())),
539 TaggedMessage::new(0, Message::User("User".to_string().into())),
540 create_tool_result_message(1, "test_tool"),
541 create_text_assistant_message(1, "Assistant"),
542 create_tool_result_message(2, "another_tool"),
543 ];
544 let response = create_response_with_messages(messages);
545 let results = response.tool_results().unwrap();
546 assert_eq!(results.len(), 2);
547 assert_eq!(results[0].tool.name, "test_tool");
548 assert_eq!(results[1].tool.name, "another_tool");
549 }
550
551 #[test]
552 fn test_generate_text_response_tool_results_no_tool_messages_but_others_present() {
553 let messages = vec![
554 TaggedMessage::new(0, Message::System("System".to_string().into())),
555 TaggedMessage::new(0, Message::User("User".to_string().into())),
556 create_text_assistant_message(0, "Assistant"),
557 ];
558 let response = create_response_with_messages(messages);
559 assert!(response.tool_results().is_none());
560 }
561
562 #[test]
563 fn test_generate_text_response_tool_results_duplicate_tool_entries() {
564 let messages = vec![
565 create_tool_result_message(0, "tool1"),
566 create_tool_result_message(1, "tool1"), create_tool_result_message(2, "tool1"), ];
569 let response = create_response_with_messages(messages);
570 let results = response.tool_results().unwrap();
571 assert_eq!(results.len(), 3);
572 assert_eq!(results[0].tool.name, "tool1");
573 assert_eq!(results[1].tool.name, "tool1");
574 assert_eq!(results[2].tool.name, "tool1");
575 }
576
577 #[test]
578 fn test_generate_text_response_tool_results_preserving_original_message_order() {
579 let messages = vec![
580 TaggedMessage::new(0, Message::System("System".to_string().into())),
581 create_tool_result_message(1, "tool1"),
582 TaggedMessage::new(1, Message::User("User".to_string().into())),
583 create_tool_result_message(2, "tool2"),
584 create_text_assistant_message(2, "Assistant"),
585 create_tool_result_message(3, "tool3"),
586 ];
587 let response = create_response_with_messages(messages);
588 let results = response.tool_results().unwrap();
589 assert_eq!(results.len(), 3);
590 assert_eq!(results[0].tool.name, "tool1");
591 assert_eq!(results[1].tool.name, "tool2");
592 assert_eq!(results[2].tool.name, "tool3");
593 }
594
595 #[test]
596 fn test_generate_text_response_tool_results_large_number_of_messages() {
597 let mut messages = Vec::new();
598 for i in 0..1000 {
600 messages.push(create_tool_result_message(0, &format!("tool{i}")));
601 if i % 100 == 0 {
602 messages.push(TaggedMessage::new(
603 0,
604 Message::User(format!("User message {i}").into()),
605 ));
606 }
607 }
608 let response = create_response_with_messages(messages);
609 let results = response.tool_results().unwrap();
610 assert_eq!(results.len(), 1000);
611 for (i, result) in results.iter().enumerate() {
612 assert_eq!(result.tool.name, format!("tool{i}"));
613 }
614 }
615}