Skip to main content

aisdk/core/language_model/
generate_text.rs

1//! Text Generation impl for the `LanguageModelRequest` trait.
2
3use 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    /// Generates text and executes tools using the language model.
22    ///
23    /// This method performs non-streaming text generation, potentially involving multiple
24    /// steps of tool calling and execution until the conversation reaches a natural stopping point.
25    /// The model may call tools based on the configured options, and responses are processed
26    /// iteratively until completion.
27    ///
28    /// For streaming responses, use [`stream_text`](Self::stream_text) instead.
29    ///
30    /// # Returns
31    ///
32    /// A [`GenerateTextResponse`] containing the final conversation state and generated content.
33    ///
34    /// # Errors
35    ///
36    /// Returns an [`Error`] if the underlying language model fails to generate a response
37    /// or if tool execution encounters an error.
38    ///
39    /// # Examples
40    ///
41    /// ```rust,no_run
42    ///# #[cfg(feature = "openai")]
43    ///# {
44    ///    use aisdk::{
45    ///        core::{LanguageModelRequest},
46    ///        providers::OpenAI,
47    ///    };
48    ///
49    ///    async fn main() -> Result<(), Box<dyn std::error::Error>> {
50    ///
51    ///        let openai = OpenAI::gpt_5();
52    ///
53    ///        let result = LanguageModelRequest::builder()
54    ///            .model(openai)
55    ///            .prompt("What is the meaning of life?")
56    ///            .build()
57    ///            .generate_text()
58    ///            .await?;
59    ///
60    ///        println!("{}", result.text().unwrap());
61    ///        Ok(())
62    ///    }
63    ///# }
64    /// ```
65    ///
66    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            // Update the current step
84            options.current_step_id += 1;
85
86            // Prepare the next step
87            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                        // add tool message
127                        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            // Finish the step
142            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            // Stop If
154            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// ============================================================================
175// Section: response types
176// ============================================================================
177
178/// Response from a generate call on `GenerateText`.
179#[derive(Debug, Clone)]
180pub struct GenerateTextResponse {
181    /// The options that generated this response
182    pub options: LanguageModelOptions,
183}
184
185impl GenerateTextResponse {
186    /// Deserializes the response text into a structured type.
187    ///
188    /// This method attempts to parse the generated text as JSON and deserialize it
189    /// into the specified type `T`. It requires that the response contains text content.
190    ///
191    /// # Type Parameters
192    ///
193    /// * `T` - The type to deserialize into, which must implement [`DeserializeOwned`].
194    ///
195    /// # Returns
196    ///
197    /// A result containing the deserialized value or a JSON error.
198    ///
199    /// # Errors
200    ///
201    /// Returns an error if there is no text response or if deserialization fails.
202    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    /// Returns the step ids of the messages in the response.
212    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    // Tests for GenerateTextResponse tool_calls()
396    #[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"), // Same name
465            create_tool_call_message(2, "tool1"), // Same name again
466        ];
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    // Tests for GenerateTextResponse tool_results()
494    #[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"), // Same name
567            create_tool_result_message(2, "tool1"), // Same name again
568        ];
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        // Add 1000 messages with tool results interspersed
599        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}