1use headless_lms_utils::cache::Cache;
5use std::ops::DerefMut;
6
7use futures::StreamExt;
8use futures::stream::BoxStream;
9use headless_lms_base::config::ApplicationConfiguration;
10use headless_lms_models::chatbot_conversation_messages::{self, ChatbotConversationMessage};
11use headless_lms_models::chatbot_conversation_messages_citations::{
12 self, ChatbotConversationMessageCitation,
13};
14use tracing::trace;
15use url::Url;
16
17use crate::azure_chatbot::azure::protocol::{
18 AISearchOutput, OutputItem, ReceivedOutputItem, ResponseOutput, check_response_complete,
19 check_response_output,
20};
21use crate::azure_chatbot::azure::sse::{AzureStreamEvent, ParsedResponseLine};
22use crate::azure_chatbot::azure::transport::ResponseLinesStream;
23use crate::azure_chatbot::client_tool_calls::abort::refused_call_output;
24use crate::azure_chatbot::events::{StreamItem, TurnEvent};
25use crate::azure_chatbot::request::replayable_input_message;
26use crate::chatbot_error::ChatbotResult;
27use crate::chatbot_tools::{
28 ChatbotToolCallResult, call_chatbot_tool, check_client_tool_call, tool_is_answered_by_client,
29};
30use crate::citations::chatbot_cited_documents_to_citations;
31use crate::llm_utils::{APIInputMessage, APIOutputMessage, MessageContent};
32use crate::prelude::*;
33use crate::user_context::ChatbotTurnContext;
34
35enum StoragePlan {
39 Insert(ChatbotConversationMessage),
40 InsertAndCite {
41 message: ChatbotConversationMessage,
42 document_urls: Vec<Url>,
43 response_id: String,
44 },
45}
46
47fn storage_plan(item: OutputItem, conversation_id: Uuid) -> ChatbotResult<StoragePlan> {
55 match item {
56 OutputItem::AzureAiSearchCall { .. } | OutputItem::Reasoning { .. } => {
57 let message = APIOutputMessage { message_type: item }
58 .to_chatbot_conversation_message(conversation_id)?;
59 Ok(StoragePlan::Insert(message))
60 }
61 OutputItem::AzureAiSearchCallOutput {
62 call_id,
63 output,
64 response_id,
65 } => {
66 let document_urls = match serde_json::from_str::<AISearchOutput>(&output) {
69 Ok(search_output) => search_output.get_urls,
70 Err(error) => {
71 warn!("Storing an Azure AI Search output that carries no citations: {error}");
72 Vec::new()
73 }
74 };
75 let message = APIOutputMessage {
76 message_type: OutputItem::AzureAiSearchCallOutput {
77 call_id,
78 output,
79 response_id: response_id.clone(),
80 },
81 }
82 .to_chatbot_conversation_message(conversation_id)?;
83 if document_urls.is_empty() {
84 return Ok(StoragePlan::Insert(message));
85 }
86 Ok(StoragePlan::InsertAndCite {
87 message,
88 document_urls,
89 response_id,
90 })
91 }
92 OutputItem::Message {
93 content: content @ MessageContent::Refusal(..),
94 response_id,
95 role,
96 } => {
97 let message = APIOutputMessage {
98 message_type: OutputItem::Message {
99 content,
100 response_id,
101 role,
102 },
103 }
104 .to_chatbot_conversation_message(conversation_id)?;
105 Ok(StoragePlan::Insert(message))
106 }
107 OutputItem::Message { .. } => Err(chatbot_err!(
108 UnexpectedProtocolShape,
109 "Unexpected message output item, it should have been streamed.".to_string()
110 )),
111 OutputItem::FunctionCall { .. } => Err(chatbot_err!(
112 UnexpectedProtocolShape,
113 "Unexpected function call output item, it should have been processed.".to_string()
114 )),
115 OutputItem::FunctionCallOutput { .. } => Err(chatbot_err!(
116 StreamInvariantViolation,
117 "Unexpected function call output item, this shouldn't happen.".to_string()
118 )),
119 }
120}
121
122async fn store_search_output_with_citations(
128 conn: &mut PgConnection,
129 message: ChatbotConversationMessage,
130 document_urls: Vec<Url>,
131 response_id: &str,
132 conversation_id: Uuid,
133 app_config: &ApplicationConfiguration,
134) -> ChatbotResult<ChatbotConversationMessage> {
135 let api_key = if let Some(azure_config) = &app_config.azure_configuration
136 && let Some(search_config) = &azure_config.search_config
137 {
138 &search_config.search_api_key
139 } else {
140 return Err(chatbot_err!(
141 Other,
142 "Azure search configuration not found, cannot process Azure AI search output item."
143 .to_string()
144 ));
145 };
146
147 let conversation_message = chatbot_conversation_messages::insert(conn, message).await?;
148
149 let res = chatbot_cited_documents_to_citations(
150 conn,
151 app_config.test_chatbot,
152 document_urls,
153 api_key,
154 conversation_message.id,
155 conversation_id,
156 )
157 .await;
158
159 if let Err(e) = res {
160 error!("Failed to save cited documents in the DB. Response id: {response_id} Error: {e}");
161 };
162
163 Ok(conversation_message)
164}
165
166pub(super) async fn store_output_item(
172 conn: &mut PgConnection,
173 item: OutputItem,
174 conversation_id: Uuid,
175 app_config: &ApplicationConfiguration,
176) -> ChatbotResult<ChatbotConversationMessage> {
177 match storage_plan(item, conversation_id)? {
178 StoragePlan::Insert(message) => {
179 Ok(chatbot_conversation_messages::insert(conn, message).await?)
180 }
181 StoragePlan::InsertAndCite {
182 message,
183 document_urls,
184 response_id,
185 } => {
186 store_search_output_with_citations(
187 conn,
188 message,
189 document_urls,
190 &response_id,
191 conversation_id,
192 app_config,
193 )
194 .await
195 }
196 }
197}
198
199pub(super) fn is_stored_by_round(item: &OutputItem) -> bool {
205 match item {
206 OutputItem::FunctionCall { .. } | OutputItem::FunctionCallOutput { .. } => true,
207 OutputItem::Message { .. }
208 | OutputItem::Reasoning { .. }
209 | OutputItem::AzureAiSearchCall { .. }
210 | OutputItem::AzureAiSearchCallOutput { .. } => false,
211 }
212}
213
214enum PendingRoundItem {
221 FunctionCall {
222 tool_name: String,
223 call_id: String,
224 arguments: String,
225 },
226 Passthrough(OutputItem),
227}
228
229struct ToolRound {
233 pending_items: Vec<PendingRoundItem>,
234 next_round_input: Vec<APIInputMessage>,
235 response_id: String,
237 suspended: bool,
238}
239
240impl ToolRound {
241 fn new(response_id: String) -> Self {
242 Self {
243 pending_items: Vec::new(),
244 next_round_input: Vec::new(),
245 response_id,
246 suspended: false,
247 }
248 }
249
250 fn queue(&mut self, item: OutputItem) -> ChatbotResult<()> {
258 let pending = match item {
259 OutputItem::FunctionCall {
260 tool_name,
261 call_id,
262 arguments,
263 ..
264 } => PendingRoundItem::FunctionCall {
265 tool_name,
266 call_id,
267 arguments,
268 },
269 OutputItem::Reasoning { .. }
270 | OutputItem::AzureAiSearchCall { .. }
271 | OutputItem::AzureAiSearchCallOutput { .. } => PendingRoundItem::Passthrough(item),
272 OutputItem::Message { .. } | OutputItem::FunctionCallOutput { .. } => {
273 return Err(chatbot_err!(
274 UnexpectedProtocolShape,
275 "Unexpected output item queued for the round's finalize pass.".to_string()
276 ));
277 }
278 };
279 self.pending_items.push(pending);
280 Ok(())
281 }
282
283 fn has_function_calls(&self) -> bool {
284 self.pending_items
285 .iter()
286 .any(|item| matches!(item, PendingRoundItem::FunctionCall { .. }))
287 }
288}
289
290enum PlannedToolCall {
292 Suspend,
294 Run,
296 Refuse(String),
298}
299
300async fn plan_tool_call(
308 conn: &mut PgConnection,
309 user_context: &ChatbotTurnContext,
310 tool_name: &str,
311 arguments: &str,
312) -> ChatbotResult<PlannedToolCall> {
313 if !tool_is_answered_by_client(tool_name) {
314 return Ok(PlannedToolCall::Run);
315 }
316 match check_client_tool_call(conn, user_context, tool_name, arguments).await {
317 Ok(Ok(())) => Ok(PlannedToolCall::Suspend),
318 Ok(Err(refusal)) => Ok(PlannedToolCall::Refuse(
319 refused_call_output(refusal, tool_name).to_string(),
320 )),
321 Err(error) => Ok(PlannedToolCall::Refuse(recover_or_terminate(
322 error,
323 tool_name,
324 "A client chatbot tool call was refused before the turn could suspend on it, reporting the failure to the LLM.",
325 )?)),
326 }
327}
328
329const UNSTORABLE_OUTPUT_PLACEHOLDER: &str = "The tool ran, but its result could not be stored and \
332is no longer available. Tell the user the lookup did not come back, or try again with a narrower \
333call.";
334
335async fn record_tool_call_with_fallback(
341 conn: &mut PgConnection,
342 conversation_id: Uuid,
343 response_id: &str,
344 call_id: &str,
345 tool_name: &str,
346 result: ChatbotToolCallResult,
347) -> ChatbotResult<Vec<APIInputMessage>> {
348 let arguments = result.arguments.clone();
349 let output_bytes = result.output.len();
350 let error = match record_tool_call(
351 conn,
352 conversation_id,
353 response_id,
354 call_id,
355 tool_name,
356 result,
357 )
358 .await
359 {
360 Ok(recorded) => return Ok(recorded),
361 Err(error) => error,
362 };
363 error!(
364 "Could not store the output of {tool_name} ({output_bytes} bytes). Storing a placeholder output instead. Error: {error:?}"
365 );
366 record_tool_call(
367 conn,
368 conversation_id,
369 response_id,
370 call_id,
371 tool_name,
372 ChatbotToolCallResult {
373 arguments,
374 output: UNSTORABLE_OUTPUT_PLACEHOLDER.to_string(),
375 citations: Vec::new(),
376 },
377 )
378 .await
379}
380
381async fn record_tool_call(
387 conn: &mut PgConnection,
388 conversation_id: Uuid,
389 response_id: &str,
390 call_id: &str,
391 tool_name: &str,
392 result: ChatbotToolCallResult,
393) -> ChatbotResult<Vec<APIInputMessage>> {
394 let citations = result.citations;
395 let tool_call_message = APIOutputMessage {
396 message_type: OutputItem::FunctionCall {
397 response_id: response_id.to_owned(),
398 call_id: call_id.to_owned(),
399 tool_name: tool_name.to_owned(),
400 arguments: result.arguments,
401 },
402 };
403 let output_message = APIOutputMessage {
404 message_type: OutputItem::FunctionCallOutput {
405 call_id: call_id.to_owned(),
406 output: result.output,
407 response_id: response_id.to_owned(),
408 },
409 };
410
411 let mut tx = conn.begin().await?;
412 let stored_call = chatbot_conversation_messages::insert(
413 &mut tx,
414 tool_call_message.to_chatbot_conversation_message(conversation_id)?,
415 )
416 .await?;
417 let stored_output = chatbot_conversation_messages::insert(
418 &mut tx,
419 output_message.to_chatbot_conversation_message(conversation_id)?,
420 )
421 .await?;
422
423 if !citations.is_empty() {
424 let (rows, page_ids) = citations
425 .into_iter()
426 .map(|citation| {
427 (
428 ChatbotConversationMessageCitation {
429 conversation_message_id: stored_output.id,
430 conversation_id,
431 title: citation.title,
432 content: citation.snippet,
433 document_url: citation.document_url,
434 citation_number: citation.citation_number,
435 ..Default::default()
436 },
437 Some(citation.page_id),
438 )
439 })
440 .unzip();
441 chatbot_conversation_messages_citations::insert_batch(&mut tx, rows, page_ids).await?;
442 }
443
444 tx.commit().await?;
445
446 Ok(vec![
447 APIInputMessage::try_from(stored_call)?,
448 APIInputMessage::try_from(stored_output)?,
449 ])
450}
451
452fn item_without_reasoning_payload(item: &OutputItem) -> OutputItem {
457 match item {
458 OutputItem::Reasoning {
459 response_id, id, ..
460 } => OutputItem::Reasoning {
461 response_id: response_id.clone(),
462 id: id.clone(),
463 summary: Vec::new(),
464 encrypted_content: None,
465 },
466 other => other.clone(),
467 }
468}
469
470#[allow(clippy::too_many_arguments)]
487pub(super) async fn parse_tool<'a, C>(
488 mut conn: C,
489 app_config: &'a ApplicationConfiguration,
490 cache: &'a Cache,
491 mut lines: ResponseLinesStream<'a>,
492 conversation_id: Uuid,
493 response_id: String,
494 user_context: &'a ChatbotTurnContext,
495 calls_from_classification: Vec<OutputItem>,
496) -> BoxStream<'a, ChatbotResult<TurnEvent>>
497where
498 C: DerefMut<Target = PgConnection> + Send + 'a,
499{
500 let mut round = ToolRound::new(response_id);
501 let mut response_received = false;
502 let mut response_incomplete = false;
503 let mut preceding_event: Option<AzureStreamEvent> = None;
504
505 trace!("Parsing tool calls...");
506
507 Box::pin(async_stream::try_stream! {
508 for call in calls_from_classification {
509 round.queue(call)?;
510 }
511 while let Some(val) = lines.next().await {
512 let line = val?;
513 let response_output: ResponseOutput = match ParsedResponseLine::parse(&line)? {
514 Some(ParsedResponseLine::Event(event)) => {
515 trace!("Event: {event:?}");
516 match &event {
517 AzureStreamEvent::ResponseCompleted => {
518 response_received = true;
519 }
520 AzureStreamEvent::Incomplete => {
521 response_received = true;
522 response_incomplete = true;
523 }
524 AzureStreamEvent::OutputTextDelta => {
525 Err(chatbot_err!(UnexpectedProtocolShape,
526 "Error: Received response text while parsing tool calls. Either the tool call parsing failed or the LLM responded in an unexpected way."
527 ))?
528 }
529 AzureStreamEvent::ErrorReported => {
530 }
532 _ => {}
533 };
534 preceding_event = Some(event);
535 continue;
536 }
537 Some(ParsedResponseLine::Data(data)) => *data,
538 None => {
539 continue;
540 }
541 };
542
543 let event = AzureStreamEvent::of_data_line(preceding_event.take(), response_output.response_type.as_deref());
544
545 check_response_output(&response_output, Some(&round.response_id), "streaming_tool_call_round")?;
546
547 if response_received {
548 check_response_complete(&response_output, response_incomplete)?;
551 if !round.has_function_calls() {
552 Err(chatbot_err!(StreamInvariantViolation,
553 "The LLM response was supposed to contain function calls, but no function calls were found"
554 ))?
555 }
556 let response_id = round.response_id.clone();
557
558 for pending_item in std::mem::take(&mut round.pending_items) {
559 let (name, id, args) = match pending_item {
560 PendingRoundItem::FunctionCall { tool_name, call_id, arguments } => {
561 (tool_name, call_id, arguments)
562 }
563 PendingRoundItem::Passthrough(item) => {
567 let stored = store_output_item(&mut conn, item, conversation_id, app_config).await?;
568 if let Some(input) = replayable_input_message(stored)? {
569 round.next_round_input.push(input);
570 }
571 continue;
572 }
573 };
574 let refused_client_call = match plan_tool_call(&mut conn, user_context, &name, &args).await? {
575 PlannedToolCall::Suspend => {
576 let tool_call_message = APIOutputMessage {
581 message_type: OutputItem::FunctionCall {
582 response_id: response_id.clone(),
583 call_id: id,
584 tool_name: name,
585 arguments: args,
586 },
587 };
588 chatbot_conversation_messages::insert(
589 &mut conn,
590 tool_call_message.to_chatbot_conversation_message(conversation_id)?,
591 )
592 .await?;
593 round.suspended = true;
594 continue;
595 }
596 PlannedToolCall::Refuse(output) => Some(output),
597 PlannedToolCall::Run => None,
598 };
599
600 let tool_result = if let Some(output) = refused_client_call {
601 ChatbotToolCallResult {
602 arguments: args,
603 output,
604 citations: Vec::new(),
605 }
606 } else {
607 let tool_call =
611 call_chatbot_tool(&mut conn, app_config, cache, &name, &args, user_context).await;
612 match tool_call {
613 Ok(result) => result,
614 Err(error) => ChatbotToolCallResult {
615 output: recover_or_terminate(
616 error,
617 &name,
618 "Chatbot tool call failed, reporting the failure to the LLM.",
619 )?,
620 arguments: args,
621 citations: Vec::new(),
622 },
623 }
624 };
625
626 let recorded = record_tool_call_with_fallback(
627 &mut conn,
628 conversation_id,
629 &response_id,
630 &id,
631 &name,
632 tool_result,
633 )
634 .await?;
635 round.next_round_input.extend(recorded);
636
637 yield TurnEvent::Item(StreamItem::ServerToolOutput { call_id: id });
638 }
639
640 if round.suspended {
641 yield TurnEvent::Suspended;
644 } else {
645 yield TurnEvent::Messages(std::mem::take(&mut round.next_round_input));
646 }
647 return;
648 } else if let Some(item) = response_output.item.and_then(ReceivedOutputItem::known) {
649 let finished = matches!(event, Some(AzureStreamEvent::OutputItemDone));
650 match &item {
651 OutputItem::FunctionCall { tool_name, call_id, arguments, .. } => {
652 if finished {
656 round.pending_items.push(PendingRoundItem::FunctionCall {
657 tool_name: tool_name.clone(),
658 call_id: call_id.clone(),
659 arguments: arguments.clone(),
660 });
661 }
662 yield TurnEvent::Item(StreamItem::Received { item, finished: false });
663 }
664 OutputItem::Message { .. } if !finished => {}
667 OutputItem::Message { content, .. } => {
668 if let MessageContent::Refusal(..) = content {
669 let stored = store_output_item(&mut conn, item, conversation_id, app_config).await?;
672 let message_id = stored.id;
673 let text = match &stored.message {
674 chatbot_conversation_messages::Message::Text(text_message) => {
675 text_message.text.clone()
676 }
677 other => Err(chatbot_err!(
678 StreamInvariantViolation,
679 format!("A stored refusal message came back as {other:?}.")
680 ))?,
681 };
682 round.next_round_input.push(APIInputMessage::try_from(stored)?);
683 yield TurnEvent::Refusal { text, message_id };
684 } else {
685 Err(chatbot_err!(
686 UnexpectedProtocolShape,
687 "Received a message item while parsing tool calls.".to_string()
688 ))?}
689 },
690 _ => {
691 if finished {
696 yield TurnEvent::ItemAnnounced(item_without_reasoning_payload(&item));
697 round.queue(item)?;
698 } else {
699 yield TurnEvent::Item(StreamItem::Received { item, finished });
700 }
701 }
702 }
703 }
704 }
705 Err(chatbot_err!(StreamEndedEarly, "Stream ended unexpectedly"))?;
709 })
710}
711
712fn recover_or_terminate(
717 error: ChatbotError,
718 tool_name: &str,
719 context: &str,
720) -> ChatbotResult<String> {
721 if error.error_type().should_terminate_stream() {
722 return Err(error);
723 }
724 warn!("{context} Tool: {tool_name}. Error: {error:?}");
725 Ok(tool_failure_output_for_llm(&error))
726}
727
728fn tool_failure_output_for_llm(error: &ChatbotError) -> String {
735 let message = match error.error_type() {
736 ChatbotErrorType::FailedAzureResponse => "Azure response failed",
737 ChatbotErrorType::AzureAISearchFilterError => "Couldn't create search filter for AI search",
738 _ => error.message(),
739 };
740 format!(
741 "The tool call failed and returned no data. Message: {message} Answer the user without this tool, or tell them what you would need to answer."
742 )
743}
744
745#[cfg(test)]
746mod tests {
747 use headless_lms_models::{
748 insert_data,
749 test_helper::{Conn, insert_chatbot_conversation},
750 };
751
752 use super::*;
753 use crate::azure_chatbot::azure::protocol::InputItem;
754 use crate::azure_chatbot::test_helpers::{azure_response_stream, shape};
755 use crate::chatbot_tools::tool_authorization::test_helpers::context;
756
757 #[tokio::test]
762 async fn only_the_finished_copy_of_a_streamed_item_reaches_the_next_round() {
763 insert_data!(:tx);
764 let (_configuration, conversation_id) = insert_chatbot_conversation(tx.as_mut()).await;
765 let user_context = context(None, None, Vec::new());
766 let app_config =
767 ApplicationConfiguration::mock_conf().expect("the mock configuration builds");
768 let cache = Cache::new("redis://127.0.0.1:1").expect("cache");
769
770 let mut events = parse_tool(
771 tx.as_mut() as &mut PgConnection,
772 &app_config,
773 &cache,
774 azure_response_stream(&[
775 "event: response.output_item.added",
776 r#"data: {"type":"response.output_item.added","item":{"type":"reasoning","id":"rs_1","response_id":"resp_1","summary":[]}}"#,
777 "event: response.output_item.done",
778 r#"data: {"type":"response.output_item.done","item":{"type":"reasoning","id":"rs_1","response_id":"resp_1","summary":[],"encrypted_content":"payload"}}"#,
779 "event: response.output_item.done",
780 r#"data: {"type":"response.output_item.done","item":{"type":"function_call","id":"fc_1","response_id":"resp_1","call_id":"call_1","name":"no_such_tool","arguments":"{}"}}"#,
781 "event: response.completed",
782 r#"data: {"type":"response.completed","response":{"id":"resp_1"}}"#,
783 ]),
784 conversation_id,
785 "resp_1".to_string(),
786 &user_context,
787 Vec::new(),
788 )
789 .await;
790
791 let mut next_round = None;
792 while let Some(event) = events.next().await {
793 if let TurnEvent::Messages(messages) = event.expect("the round streams to the end") {
794 next_round = Some(messages);
795 }
796 }
797
798 let next_round = next_round.expect("the round hands its items on");
799 assert_eq!(
800 shape(&next_round),
801 vec!["reasoning:rs_1", "call:call_1", "output:call_1"],
802 );
803 let InputItem::Reasoning {
804 encrypted_content, ..
805 } = &next_round[0].message_type
806 else {
807 panic!("the first item is the reasoning item");
808 };
809 assert_eq!(encrypted_content.as_deref(), Some("payload"));
810 }
811
812 #[test]
816 fn a_non_conforming_search_output_is_stored_without_citations() {
817 let item = OutputItem::AzureAiSearchCallOutput {
818 response_id: "resp_1".to_string(),
819 call_id: "call_1".to_string(),
820 output: "remote tool call failed".to_string(),
821 };
822
823 let plan =
824 storage_plan(item, Uuid::new_v4()).expect("a non-conforming output still stores");
825 assert!(matches!(plan, StoragePlan::Insert(_)));
826 }
827
828 #[tokio::test]
831 async fn a_stream_that_ends_before_response_completed_errors() {
832 insert_data!(:tx);
833 let (_configuration, conversation_id) = insert_chatbot_conversation(tx.as_mut()).await;
834 let user_context = context(None, None, Vec::new());
835 let app_config =
836 ApplicationConfiguration::mock_conf().expect("the mock configuration builds");
837 let cache = Cache::new("redis://127.0.0.1:1").expect("cache");
838
839 let mut events = parse_tool(
840 tx.as_mut() as &mut PgConnection,
841 &app_config,
842 &cache,
843 azure_response_stream(&[
844 "event: response.output_item.done",
845 r#"data: {"type":"response.output_item.done","item":{"type":"function_call","id":"fc_1","response_id":"resp_1","call_id":"call_1","name":"no_such_tool","arguments":"{}"}}"#,
846 ]),
847 conversation_id,
848 "resp_1".to_string(),
849 &user_context,
850 Vec::new(),
851 )
852 .await;
853
854 let error = loop {
855 match events.next().await.expect("the stream ends in an error") {
856 Ok(_) => continue,
857 Err(error) => break error,
858 }
859 };
860 assert_eq!(*error.error_type(), ChatbotErrorType::StreamEndedEarly);
861 }
862}