Skip to main content

headless_lms_chatbot/azure_chatbot/turn/
mod.rs

1//! The turn driver: keeps asking Azure while a round answers itself with tool calls, and ends on
2//! an answer, an error, a suspension, or the round budget.
3
4mod 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
48/// How many LLM requests one turn may make, bounding a model that keeps calling tools instead of
49/// answering.
50const MAX_TOOL_CALL_ROUNDS_PER_TURN: u32 = 15;
51
52/// Starts a turn for a new user message, and streams its NDJSON events to the client.
53#[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/// Records a client's answer to a tool call the turn suspended on, and continues that turn once
80/// nothing else is outstanding.
81///
82/// Of a round of parallel calls, only the request that answers the last one gets the resumed turn;
83/// the others get a stream carrying `Suspended` again, so a client reads every response the same
84/// way. `tool_call_id` must be a client-answered call of `conversation_id` that has no answer yet
85/// and `answer` must fit what that call offered, or this fails with
86/// [ChatbotErrorType::InvalidToolAnswer] and writes nothing.
87#[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
113/// What starts a turn: a new message from the user, or one resuming after the client answered a
114/// tool call the previous turn suspended on.
115enum 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
128/// Shared preamble of both ways a turn can begin: acquire a connection, repair any tool call a
129/// dead turn of this conversation left unanswered, build the request the turn runs with, and hand
130/// off to [stream_turn] — or, on a resume that is still waiting on another call, return the
131/// [ChatbotChatStreamEvent::Suspended] stream without ever reaching it.
132///
133/// Repairing before either path reads the conversation's history is required, not incidental: an
134/// unanswered call from a dead turn makes the LLM reject every later message of the conversation.
135/// Only long-dead calls are touched: another request may be streaming a turn of this same
136/// conversation.
137async 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            // One transaction for the answer path: a confirmed action tool's mutation, its audit
172            // row (both inside `client_tool_output_for_answer`), and the recorded tool output
173            // below commit or roll back together, so the transcript can never claim an effect the
174            // database does not have.
175            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                // Another client tool call is still open, so there's no resumed turn to ride the
205                // payload ahead of -- emit it here or it's lost, and the browser never sees it.
206                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            // The reset link (or similar) an executed action tool produced is for this browser
234            // only and is never persisted, so it can only reach the client by riding ahead of the
235            // resumed turn's own stream.
236            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
268/// What a round that ended in error becomes: logs it, answers whatever tool call the turn left
269/// without an output, and either returns the wire event for an error the turn survives or the
270/// original error for one that ends it.
271///
272/// The reap belongs here rather than at each failing site, so that no error path can end a turn
273/// without it: a call with no output makes the LLM reject every later message of the conversation.
274/// `response_ids` are the responses this turn's rounds were given, which keeps the reap off the
275/// calls of a turn streaming in another request.
276///
277/// Takes the pool rather than a connection: the call sites hold their round's connection under a
278/// live borrow, or have already given theirs back, so this acquires its own for the reap.
279async 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
316/// Records a failure that killed a turn in `error_variants`/`error_occurrences`.
317///
318/// A streaming response has already sent its headers by the time a round can fail, so actix never
319/// calls `ResponseError::error_response` for it and the backend's usual reporting path sees none
320/// of these -- without this they exist only as a log line. Best effort: the turn is already being
321/// torn down, so a failure to report is logged and otherwise ignored.
322async 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
342/// What every chatbot stream failure records beyond its message, so one query finds them all and
343/// each row says which turn it belongs to.
344///
345/// `input` is the item chain the failed round was sent. It is the only thing that names the tool
346/// call a failure belongs to once that call's row has died with the transaction that failed to
347/// store it.
348fn 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
363/// Builds the wire event for a round-ending error, folding in the input summary and response ids
364/// every call site of [recover_from_round_error] otherwise repeats.
365async 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
384/// Runs the request rounds of one turn against the LLM and streams its events as NDJSON.
385///
386/// Keeps asking the LLM as long as a round ends in tool calls it answered itself, and ends the
387/// turn on a text answer, an error, a suspension, or the iteration limit. Owns the cancellation
388/// guard, so a client that disappears mid-turn still gets what arrived saved.
389fn 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    // Shared with the guard so that its cleanup answers this turn's tool calls and no other
403    // turn's.
404    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                // Not a ChatbotError, but it ends a turn as thoroughly as one and is just as
421                // invisible outside the logs, so it is recorded the same way.
422                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            // One statement, so no guard is alive across the awaits below: `?` inside `try_stream!`
463            // parks the generator rather than returning, and a guard held at one is never released.
464            response_ids.lock().await.push(received_response_id.clone());
465
466            // Acquired only now: the request and its classification need no database, and the pool
467            // is shared with the rest of the application.
468            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                        // A function call classifies the round it opens, so it arrives here rather
475                        // than in the round that runs it; hand it over to be recorded there.
476                        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            // Some only for a round that streams an answer, which is the only round whose events
501            // address a message of their own.
502            let (mut final_stream, text_message_id) = match typed_response_stream {
503                ResponseStreamType::ToolCall(stream) => {
504                    // The round writes a row per call as it goes, so it keeps the connection.
505                    (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                    // One statement, so the guard is not alive across the awaits below: `?` inside
524                    // `try_stream!` parks the generator rather than returning, and a guard held at
525                    // one is never released.
526                    *response_message_id.lock().await = Some(response_message.id);
527
528                    // Move the citations of the turn onto the message that cites them before its
529                    // text reaches the learner, so the markers in it have something behind them.
530                    models::chatbot_conversation_messages_citations::attach_turn_citations_to_message(
531                        &mut conn,
532                        conversation_id,
533                        response_message.id,
534                    ).await?;
535
536                    // Given back before the answer streams, which takes as long as the model takes
537                    // to write it; what little the loop below stores acquires its own.
538                    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                                // Nothing ever reached this message, and an empty one that is
553                                // never completed replays into every later turn. Cleared first so
554                                // the cancellation guard does not try to clean it up again.
555                                *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                        // A `Message` among these is an unexpected wire shape, not a normal
580                        // path; store_output_item errors on it rather than storing it.
581                        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                            // A search output stored after the answer's message was created keeps
587                            // its citations on the tool-output row, out of reach of the markers in
588                            // the answer that cite them.
589                            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                        // Stored by the round that produced it, at its position among that
606                        // round's other items (see `TurnEvent::ItemAnnounced`); this only converts
607                        // and forwards it to the client.
608                        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                        // The turn ended on purpose, so the guard must not treat the conversation
636                        // as one that died mid-answer and clean up after it.
637                        done.store(true, atomic::Ordering::Relaxed);
638                        break 'outer;
639                    }
640                }
641            }
642        }
643    };
644
645    Box::pin(GuardedStream::new(guard, response_stream))
646}