1mod cancellation;
5mod round;
6mod text_response;
7
8use headless_lms_utils::cache::Cache;
9use std::pin::Pin;
10use std::sync::{
11 Arc,
12 atomic::{self, AtomicBool},
13};
14
15use bytes::Bytes;
16use futures::{Stream, StreamExt};
17use headless_lms_base::config::ApplicationConfiguration;
18use headless_lms_models::chatbot_conversation_message_messages::{
19 ChatbotConversationMessageMessage, MessageRole,
20};
21use headless_lms_models::chatbot_conversation_messages::{ChatbotConversationMessage, Message};
22use sqlx::PgPool;
23use tokio::sync::Mutex;
24use tracing::trace;
25
26use super::azure::protocol::{LLMRequest, OutputItem};
27use super::azure::sse::detect_response_kind;
28use super::azure::transport::{ResponseStreamType, make_request_and_create_stream};
29use super::client_tool_calls::answer::{client_tool_output_for_answer, rejected_tool_answer_error};
30use super::client_tool_calls::repair::{
31 answer_stale_unfinished_tool_calls, answer_unfinished_tool_calls,
32};
33use super::events::{
34 ChatbotChatStreamEvent, StreamItem, TurnEvent, error_event_from_error, error_event_from_text,
35 ndjson_line, single_event_stream, stream_event_for,
36};
37use super::request::replayable_input_message;
38use crate::chatbot_error::ChatbotResult;
39use crate::chatbot_tools::ClientToolAnswer;
40use crate::conversation_context::ChatbotPageContext;
41use crate::llm_utils::{estimate_tokens, summarize_input_for_log};
42use crate::prelude::*;
43use crate::user_context::ChatbotTurnContext;
44use cancellation::{GuardedStream, RequestCancelledGuard, save_partial_answer};
45use round::{is_stored_by_round, parse_tool, store_output_item};
46use text_response::parse_text_response;
47
48const MAX_TOOL_CALL_ROUNDS_PER_TURN: u32 = 15;
51
52#[allow(clippy::too_many_arguments)]
54pub async fn send_chat_request_and_parse_stream(
55 pool: PgPool,
56 app_configuration: &ApplicationConfiguration,
57 cache: Cache,
58 chatbot_configuration_id: Uuid,
59 conversation_id: Uuid,
60 message: &str,
61 page_context: Option<ChatbotPageContext>,
62 user_context: ChatbotTurnContext,
63) -> ChatbotResult<Pin<Box<dyn Stream<Item = ChatbotResult<Bytes>> + Send>>> {
64 begin_turn(
65 pool,
66 app_configuration,
67 cache,
68 conversation_id,
69 user_context,
70 TurnStart::NewUserMessage {
71 chatbot_configuration_id,
72 message,
73 page_context,
74 },
75 )
76 .await
77}
78
79#[allow(clippy::too_many_arguments)]
88pub async fn answer_tool_call_and_resume_stream(
89 pool: PgPool,
90 app_configuration: &ApplicationConfiguration,
91 cache: Cache,
92 chatbot_configuration_id: Uuid,
93 conversation_id: Uuid,
94 tool_call_id: &str,
95 answer: &ClientToolAnswer,
96 user_context: ChatbotTurnContext,
97) -> ChatbotResult<Pin<Box<dyn Stream<Item = ChatbotResult<Bytes>> + Send>>> {
98 begin_turn(
99 pool,
100 app_configuration,
101 cache,
102 conversation_id,
103 user_context,
104 TurnStart::ResumedFromToolAnswer {
105 chatbot_configuration_id,
106 tool_call_id,
107 answer,
108 },
109 )
110 .await
111}
112
113enum TurnStart<'a> {
116 NewUserMessage {
117 chatbot_configuration_id: Uuid,
118 message: &'a str,
119 page_context: Option<ChatbotPageContext>,
120 },
121 ResumedFromToolAnswer {
122 chatbot_configuration_id: Uuid,
123 tool_call_id: &'a str,
124 answer: &'a ClientToolAnswer,
125 },
126}
127
128async fn begin_turn(
138 pool: PgPool,
139 app_configuration: &ApplicationConfiguration,
140 cache: Cache,
141 conversation_id: Uuid,
142 user_context: ChatbotTurnContext,
143 start: TurnStart<'_>,
144) -> ChatbotResult<Pin<Box<dyn Stream<Item = ChatbotResult<Bytes>> + Send>>> {
145 let mut conn = pool.acquire().await?;
146 let unanswered = answer_stale_unfinished_tool_calls(&mut conn, conversation_id).await?;
147 let app_config = app_configuration.to_owned();
148
149 let chat_request = match start {
150 TurnStart::NewUserMessage {
151 chatbot_configuration_id,
152 message,
153 page_context,
154 } => {
155 LLMRequest::build_and_insert_incoming_user_message_to_db(
156 &mut conn,
157 chatbot_configuration_id,
158 conversation_id,
159 message,
160 page_context,
161 &user_context,
162 &app_config,
163 )
164 .await?
165 }
166 TurnStart::ResumedFromToolAnswer {
167 chatbot_configuration_id,
168 tool_call_id,
169 answer,
170 } => {
171 let mut tx = conn.begin().await?;
176
177 let answered = client_tool_output_for_answer(
178 &mut tx,
179 &app_config,
180 conversation_id,
181 &unanswered,
182 tool_call_id,
183 answer,
184 &user_context,
185 )
186 .await?;
187
188 let outcome = models::chatbot_conversation_messages::answer_client_tool_call(
189 &mut tx,
190 conversation_id,
191 tool_call_id,
192 answered.output,
193 answered.client_answer,
194 )
195 .await
196 .map_err(rejected_tool_answer_error)?;
197
198 tx.commit().await?;
199
200 if !outcome.turn_can_resume {
201 trace!(
202 "Tool call {tool_call_id} answered, the turn is still waiting for another answer"
203 );
204 if let Some(payload) = answered.execution_payload {
207 let event = ChatbotChatStreamEvent::ActionExecuted {
208 tool_call_id: tool_call_id.to_string(),
209 payload,
210 };
211 let action_line = ndjson_line(&event)?;
212 let suspended_line = ndjson_line(&ChatbotChatStreamEvent::Suspended)?;
213 return Ok(Box::pin(futures::stream::iter([
214 Ok(action_line),
215 Ok(suspended_line),
216 ])));
217 }
218 return single_event_stream(ChatbotChatStreamEvent::Suspended);
219 }
220
221 let configuration =
222 models::chatbot_configurations::get_by_id(&mut conn, chatbot_configuration_id)
223 .await?;
224 let chat_request = LLMRequest::build_from_conversation(
225 &mut conn,
226 &configuration,
227 conversation_id,
228 &user_context,
229 &app_config,
230 )
231 .await?;
232
233 if let Some(payload) = answered.execution_payload {
237 let event = ChatbotChatStreamEvent::ActionExecuted {
238 tool_call_id: tool_call_id.to_string(),
239 payload,
240 };
241 let line = ndjson_line(&event)?;
242 return Ok(Box::pin(
243 futures::stream::once(async move { Ok(line) }).chain(stream_turn(
244 pool,
245 app_config,
246 cache,
247 conversation_id,
248 chat_request,
249 user_context,
250 )),
251 ));
252 }
253
254 chat_request
255 }
256 };
257
258 Ok(stream_turn(
259 pool,
260 app_config,
261 cache,
262 conversation_id,
263 chat_request,
264 user_context,
265 ))
266}
267
268async fn recover_from_round_error(
280 pool: &PgPool,
281 conversation_id: Uuid,
282 response_ids: &[String],
283 input_summary: &str,
284 error: ChatbotError,
285) -> ChatbotResult<Bytes> {
286 let response_id = response_ids.last().map(String::as_str);
287 error!(
288 input = %input_summary,
289 "Stream ended unexpectedly. Response id: {} Error: {}", response_id.unwrap_or("not received"), error
290 );
291 let mut conn = pool.acquire().await?;
292 report_stream_failure(
293 &mut conn,
294 error.message().to_string(),
295 Some(format!("{error:?}")),
296 stream_failure_details(
297 &format!("{:?}", error.error_type()),
298 conversation_id,
299 response_id,
300 input_summary,
301 ),
302 )
303 .await;
304 if let Err(e2) = answer_unfinished_tool_calls(&mut conn, conversation_id, response_ids).await {
305 error!(
306 "Error in chatbot streaming and couldn't answer unfinished tool calls: {e2}. Response id: {}",
307 response_id.unwrap_or("not received")
308 );
309 }
310 if error.error_type().should_terminate_stream() {
311 return Err(error);
312 }
313 error_event_from_error(&error)
314}
315
316async fn report_stream_failure(
323 conn: &mut PgConnection,
324 message: String,
325 stack_trace: Option<String>,
326 details: serde_json::Value,
327) {
328 let report = models::errors::NewErrorReport {
329 service: "headless-lms".to_string(),
330 error_source: Some(models::errors::ErrorSource::Backend),
331 message,
332 stack_trace,
333 path: None,
334 app_version: None,
335 details: Some(details),
336 };
337 if let Err(e) = models::errors::insert(conn, None, &report).await {
338 warn!("Could not record the chatbot stream failure: {e}");
339 }
340}
341
342fn stream_failure_details(
349 kind: &str,
350 conversation_id: Uuid,
351 response_id: Option<&str>,
352 input_summary: &str,
353) -> serde_json::Value {
354 serde_json::json!({
355 "kind": "chatbot_stream_error",
356 "chatbot_error_type": kind,
357 "conversation_id": conversation_id,
358 "response_id": response_id,
359 "input": input_summary,
360 })
361}
362
363async fn recover_and_summarize(
366 pool: &PgPool,
367 conversation_id: Uuid,
368 response_ids: &Mutex<Vec<String>>,
369 input: &[crate::llm_utils::APIInputMessage],
370 error: ChatbotError,
371) -> ChatbotResult<Bytes> {
372 let input_summary = summarize_input_for_log(input);
373 let round_response_ids = response_ids.lock().await.clone();
374 recover_from_round_error(
375 pool,
376 conversation_id,
377 &round_response_ids,
378 &input_summary,
379 error,
380 )
381 .await
382}
383
384fn stream_turn(
390 pool: PgPool,
391 app_config: ApplicationConfiguration,
392 cache: Cache,
393 conversation_id: Uuid,
394 mut chat_request: LLMRequest,
395 user_context: ChatbotTurnContext,
396) -> Pin<Box<dyn Stream<Item = ChatbotResult<Bytes>> + Send>> {
397 let mut rounds_left = MAX_TOOL_CALL_ROUNDS_PER_TURN;
398
399 let done = Arc::new(AtomicBool::new(false));
400 let full_response_text = Arc::new(Mutex::new(String::new()));
401 let response_message_id: Arc<Mutex<Option<Uuid>>> = Arc::new(Mutex::new(None));
402 let response_ids: Arc<Mutex<Vec<String>>> = Arc::new(Mutex::new(Vec::new()));
405
406 let guard = RequestCancelledGuard {
407 conversation_id,
408 response_ids: response_ids.clone(),
409 response_message_id: response_message_id.clone(),
410 full_response_text: full_response_text.clone(),
411 pool: pool.clone(),
412 done: done.clone(),
413 };
414
415 let response_stream = async_stream::try_stream! {
416 'outer: loop {
417 if rounds_left == 0 {
418 const ROUND_LIMIT_MESSAGE: &str = "Maximum tool call iterations exceeded";
419 error!("{ROUND_LIMIT_MESSAGE}");
420 if let Ok(mut conn) = pool.acquire().await {
423 report_stream_failure(
424 &mut conn,
425 ROUND_LIMIT_MESSAGE.to_string(),
426 None,
427 stream_failure_details(
428 "RoundLimitExceeded",
429 conversation_id,
430 response_ids.lock().await.last().map(String::as_str),
431 &summarize_input_for_log(&chat_request.input),
432 ),
433 )
434 .await;
435 }
436 yield error_event_from_text("Maximum tool call iterations exceeded. The LLM may be stuck in a loop.")?;
437 done.store(true, atomic::Ordering::Relaxed);
438 break 'outer;
439 }
440 rounds_left -= 1;
441
442 let lines = match make_request_and_create_stream(&chat_request, &app_config).await {
443 Ok(val) => val,
444 Err(error) => {
445 let event = recover_and_summarize(&pool, conversation_id, &response_ids, &chat_request.input, error).await?;
446 yield event;
447 done.store(true, atomic::Ordering::Relaxed);
448 break 'outer;
449 },
450 };
451 let classified = match detect_response_kind(lines).await {
452 Ok(classified) => classified,
453 Err(e) => {
454 let event = recover_and_summarize(&pool, conversation_id, &response_ids, &chat_request.input, e).await?;
455 yield event;
456 done.store(true, atomic::Ordering::Relaxed);
457 break 'outer;
458 },
459 };
460 let received_response_id = classified.response_id;
461 let typed_response_stream = classified.stream;
462 response_ids.lock().await.push(received_response_id.clone());
465
466 let mut conn = pool.acquire().await?;
469
470 let mut calls_from_classification = Vec::new();
471 for stream_item in classified.items {
472 if let StreamItem::Received { item, finished: true } = &stream_item {
473 if is_stored_by_round(item) {
474 calls_from_classification.push(item.to_owned());
477 } else {
478 let stored = match store_output_item(&mut conn, item.to_owned(), conversation_id, &app_config)
479 .await
480 .and_then(replayable_input_message)
481 {
482 Ok(stored) => stored,
483 Err(e) => {
484 let event = recover_and_summarize(&pool, conversation_id, &response_ids, &chat_request.input, e).await?;
485 yield event;
486 done.store(true, atomic::Ordering::Relaxed);
487 break 'outer;
488 }
489 };
490 if let Some(stored) = stored {
491 chat_request.input.push(stored);
492 }
493 }
494 }
495 if let Some(event) = stream_event_for(stream_item) {
496 yield ndjson_line(&event)?;
497 };
498 }
499
500 let (mut final_stream, text_message_id) = match typed_response_stream {
503 ResponseStreamType::ToolCall(stream) => {
504 (parse_tool(conn, &app_config, &cache, stream, conversation_id, received_response_id, &user_context, calls_from_classification).await, None)
506 }
507 ResponseStreamType::TextResponse(stream) => {
508 let response_message = models::chatbot_conversation_messages::insert(
509 &mut conn,
510 ChatbotConversationMessage {
511 conversation_id,
512 message: Message::Text(ChatbotConversationMessageMessage {
513 text: "".to_string(),
514 message_role: MessageRole::Assistant,
515 message_is_complete: false,
516 response_id: Some(received_response_id.clone()),
517 ..Default::default()
518 }),
519 ..Default::default()
520 },
521 ).await?;
522
523 *response_message_id.lock().await = Some(response_message.id);
527
528 models::chatbot_conversation_messages_citations::attach_turn_citations_to_message(
531 &mut conn,
532 conversation_id,
533 response_message.id,
534 ).await?;
535
536 drop(conn);
539
540 (parse_text_response(stream, full_response_text.clone(), received_response_id).await, Some(response_message.id))
541 }
542 };
543
544 while let Some(line) = final_stream.next().await {
545 let val = match line {
546 Ok(val) => val,
547 Err(e) => {
548 if let Some(message_id) = text_message_id {
549 let full_response_as_string = full_response_text.lock().await.clone();
550 let mut conn = pool.acquire().await?;
551 if full_response_as_string.is_empty() {
552 *response_message_id.lock().await = None;
556 models::chatbot_conversation_messages::delete(&mut conn, message_id).await?;
557 } else {
558 let used_tokens = estimate_tokens(&full_response_as_string);
559 save_partial_answer(&mut conn, message_id, &full_response_as_string, used_tokens).await?;
560 }
561 };
562 let event = recover_and_summarize(&pool, conversation_id, &response_ids, &chat_request.input, e).await?;
563 yield event;
564 done.store(true, atomic::Ordering::Relaxed);
565 break 'outer;
566 }
567 };
568 match val {
569 TurnEvent::Delta(text) => {
570 match text_message_id {
571 Some(message_id) => yield ndjson_line(&ChatbotChatStreamEvent::Delta { text, message_id })?,
572 None => Err(chatbot_err!(StreamInvariantViolation, "Received answer text from a round that streams no answer."))?,
573 }
574 },
575 TurnEvent::Refusal { text, message_id } => {
576 yield ndjson_line(&ChatbotChatStreamEvent::Delta { text, message_id })?;
577 },
578 TurnEvent::Item(stream_item) => {
579 if let StreamItem::Received { item, finished: true } = &stream_item
582 && !is_stored_by_round(item)
583 {
584 let mut conn = pool.acquire().await?;
585 store_output_item(&mut conn, item.to_owned(), conversation_id, &app_config).await?;
586 if let Some(message_id) = text_message_id
590 && matches!(item, OutputItem::AzureAiSearchCallOutput { .. })
591 {
592 models::chatbot_conversation_messages_citations::attach_turn_citations_to_message(
593 &mut conn,
594 conversation_id,
595 message_id,
596 ).await?;
597 }
598 }
599
600 if let Some(response) = stream_event_for(stream_item) {
601 yield ndjson_line(&response)?;
602 };
603 },
604 TurnEvent::ItemAnnounced(item) => {
605 if let Some(response) = stream_event_for(StreamItem::Received { item, finished: true }) {
609 yield ndjson_line(&response)?;
610 };
611 },
612 TurnEvent::Messages(messages) => {
613 chat_request.input.extend(messages);
614 },
615 TurnEvent::Done { text, used_tokens } => {
616 match text_message_id {
617 Some(message_id) => {
618 let mut conn = pool.acquire().await?;
619 models::chatbot_conversation_messages::update(
620 &mut conn,
621 message_id,
622 &text,
623 true,
624 used_tokens,
625 ).await?;
626 }
627 None => Err(chatbot_err!(StreamInvariantViolation, "A round that streams no answer reported one finished."))?,
628 }
629 done.store(true, atomic::Ordering::Relaxed);
630 yield ndjson_line(&ChatbotChatStreamEvent::Done)?;
631 break 'outer;
632 }
633 TurnEvent::Suspended => {
634 yield ndjson_line(&ChatbotChatStreamEvent::Suspended)?;
635 done.store(true, atomic::Ordering::Relaxed);
638 break 'outer;
639 }
640 }
641 }
642 }
643 };
644
645 Box::pin(GuardedStream::new(guard, response_stream))
646}