Skip to main content

headless_lms_chatbot/chatbot_tools/custom_tools/
course_material_search.rs

1use headless_lms_authorization::Action;
2use headless_lms_utils::cache::Cache;
3use headless_lms_utils::course_url::build_courses_base_url;
4use std::str::FromStr;
5
6use indexmap::IndexMap;
7
8use headless_lms_models::chatbot_configurations::ToolCategory;
9use headless_lms_models::{
10    chatbot_conversation_messages_citations, courses, organizations,
11    pages::{self, PageSearchResult, SearchRequest},
12};
13use headless_lms_utils::{
14    json_schema_types::{JSONType, JsonItem, Schema, SchemaPropertyType},
15    strings::truncate_utf8_at_boundary,
16};
17
18use crate::{
19    azure_chatbot::azure::tools::{AzureLLMFunctionToolDefinition, LLMToolType},
20    chatbot_tools::{
21        ChatbotTool, ChatbotToolDeclaration, ToolCitation, ToolProperties,
22        course_scope::resolve_course_scope, tool_authorization::ToolRequirement,
23    },
24    prelude::*,
25    user_context::ChatbotTurnContext,
26};
27
28pub type CourseMaterialSearchTool = ToolProperties<CourseMaterialSearchState>;
29
30struct SearchHit {
31    page_id: Uuid,
32    title: String,
33    chapter_name: Option<String>,
34    url_path: String,
35    rank: Option<f32>,
36    snippet: Option<String>,
37    citation_number: i32,
38}
39
40#[derive(Serialize)]
41struct OutputCourse<'a> {
42    id: Uuid,
43    name: &'a str,
44    slug: &'a str,
45}
46
47#[derive(Serialize)]
48struct OutputResult<'a> {
49    page_id: Uuid,
50    title: &'a str,
51    #[serde(skip_serializing_if = "Option::is_none")]
52    chapter_name: Option<&'a str>,
53    url_path: &'a str,
54    #[serde(skip_serializing_if = "Option::is_none")]
55    rank: Option<f32>,
56    #[serde(skip_serializing_if = "Option::is_none")]
57    snippet: Option<&'a str>,
58    citation_number: i32,
59}
60
61#[derive(Serialize)]
62struct Output<'a> {
63    course: OutputCourse<'a>,
64    results: Vec<OutputResult<'a>>,
65    #[serde(skip_serializing_if = "Option::is_none")]
66    note: Option<&'static str>,
67}
68
69pub struct CourseMaterialSearchState {
70    course_id: Uuid,
71    course_name: String,
72    course_slug: String,
73    hits: Vec<SearchHit>,
74    document_url_prefix: String,
75}
76
77#[derive(Deserialize)]
78struct RawArguments {
79    course_id: String,
80    query: String,
81}
82
83pub struct CourseMaterialSearchArguments {
84    course_id: Uuid,
85    query: String,
86}
87
88/// Manual, not derived: `course_id`/`query` need validation `#[derive(Deserialize)]` can't
89/// express, and this is what [ChatbotTool::Arguments]'s `DeserializeOwned` bound is satisfied by
90/// (`parse_arguments` below is overridden and never calls it, but the bound still has to hold).
91impl<'de> Deserialize<'de> for CourseMaterialSearchArguments {
92    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
93    where
94        D: serde::Deserializer<'de>,
95    {
96        let raw = RawArguments::deserialize(deserializer)?;
97        build_arguments(raw).map_err(serde::de::Error::custom)
98    }
99}
100
101fn build_arguments(raw: RawArguments) -> ChatbotResult<CourseMaterialSearchArguments> {
102    let course_id = Uuid::from_str(&raw.course_id).map_err(|e| {
103        chatbot_err!(
104            InvalidToolArguments,
105            format!("'{}' is not a valid course_id.", raw.course_id),
106            e
107        )
108    })?;
109    let query = raw.query.trim().to_string();
110    if query.is_empty() {
111        return Err(chatbot_err!(
112            InvalidToolArguments,
113            "query must not be empty.".to_string()
114        ));
115    }
116    if query.chars().count() > MAX_QUERY_LENGTH {
117        return Err(chatbot_err!(
118            InvalidToolArguments,
119            format!(
120                "query is too long ({} characters); keep it under {MAX_QUERY_LENGTH} characters, closer to a few keywords than a paragraph.",
121                query.chars().count()
122            )
123        ));
124    }
125    Ok(CourseMaterialSearchArguments { course_id, query })
126}
127
128const MAX_RESULTS: usize = 10;
129const MAX_QUERY_LENGTH: usize = 200;
130
131/// Marks from `ts_headline` on the raw match, useless once handed to the model.
132fn strip_headline_marks(headline: Option<String>) -> Option<String> {
133    headline.map(|s| s.replace("<b>", "").replace("</b>", ""))
134}
135
136impl ChatbotToolDeclaration for CourseMaterialSearchTool {
137    const NAME: &'static str = "course_material_search";
138
139    fn offer_requirements(user_context: &ChatbotTurnContext) -> Vec<ToolRequirement> {
140        vec![ToolRequirement::on_turn(
141            Action::ViewInternalCourseStructure,
142            user_context,
143        )]
144    }
145
146    const CATEGORY: ToolCategory = ToolCategory::CourseMaterial;
147
148    fn get_tool_definition() -> AzureLLMFunctionToolDefinition {
149        AzureLLMFunctionToolDefinition {
150            tool_type: LLMToolType::Function,
151            name: Self::NAME.to_string(),
152            description: "Search a course's own pages by keyword, the same full-text search that backs the course material search dialog. Returns the pages that matched, each with a short snippet. Use this before document_lookup to find which page has what you need.".to_string(),
153            parameters: Schema::strict_object(
154                IndexMap::from([
155                    (
156                        "course_id".to_string(),
157                        SchemaPropertyType::Item(JsonItem {
158                            type_field: JSONType::String,
159                            description: Some(
160                                "The course to search. Resolve it with find_course first."
161                                    .to_string(),
162                            ),
163                        }),
164                    ),
165                    (
166                        "query".to_string(),
167                        SchemaPropertyType::Item(JsonItem {
168                            type_field: JSONType::String,
169                            description: Some(
170                                "What to look for, in the course's own language and wording. This is a keyword search, not a semantic one: prefer the words the material would use, and try a different phrasing if nothing is found."
171                                    .to_string(),
172                            ),
173                        }),
174                    ),
175                ]),
176                None,
177            ),
178            strict: true,
179        }
180    }
181}
182
183impl ChatbotTool for CourseMaterialSearchTool {
184    type Arguments = CourseMaterialSearchArguments;
185
186    fn call_requirements(
187        arguments: &Self::Arguments,
188        _user_context: &ChatbotTurnContext,
189    ) -> Vec<ToolRequirement> {
190        vec![ToolRequirement::on_course(
191            Action::ViewInternalCourseStructure,
192            arguments.course_id,
193        )]
194    }
195
196    fn parse_arguments(args_string: String) -> ChatbotResult<Self::Arguments> {
197        let raw: RawArguments = serde_json::from_str(&args_string).map_err(|e| {
198            chatbot_err!(
199                InvalidToolArguments,
200                format!("Couldn't parse tool arguments. Arguments: {args_string}"),
201                e
202            )
203        })?;
204        build_arguments(raw)
205    }
206
207    async fn from_db_and_arguments(
208        conn: &mut PgConnection,
209        app_config: &ApplicationConfiguration,
210        _cache: &Cache,
211        arguments: Self::Arguments,
212        user_context: &ChatbotTurnContext,
213    ) -> ChatbotResult<Self> {
214        let course_id = resolve_course_scope(user_context, Some(arguments.course_id))?;
215        let course = courses::get_course(conn, course_id).await.map_err(|e| {
216            chatbot_err!(
217                InvalidToolArguments,
218                format!("No course found with id {course_id}."),
219                e
220            )
221        })?;
222        let organization = organizations::get_organization(conn, course.organization_id).await?;
223
224        let search_request = SearchRequest {
225            query: arguments.query,
226        };
227        let phrase_results =
228            pages::get_page_search_results_for_phrase(conn, course_id, &search_request).await?;
229        let word_results =
230            pages::get_page_search_results_for_words(conn, course_id, &search_request).await?;
231
232        let mut merged: Vec<PageSearchResult> = phrase_results;
233        let already_present: std::collections::HashSet<Uuid> =
234            merged.iter().map(|r| r.id).collect();
235        merged.extend(
236            word_results
237                .into_iter()
238                .filter(|r| !already_present.contains(&r.id)),
239        );
240        merged.truncate(MAX_RESULTS);
241
242        // A whole turn's citations end up on one message, so numbering has to continue past
243        // whatever an earlier search already used in this turn rather than restart at zero.
244        let starting_number = if let Some(conversation_id) = user_context.conversation_id {
245            chatbot_conversation_messages_citations::max_citation_number_in_turn(
246                conn,
247                conversation_id,
248            )
249            .await?
250            .unwrap_or(0)
251        } else {
252            0
253        };
254
255        let ids_missing_headline: Vec<Uuid> = merged
256            .iter()
257            .filter(|r| r.title_headline.is_none())
258            .map(|r| r.id)
259            .collect();
260        let fallback_titles = pages::get_titles_by_ids(conn, &ids_missing_headline).await?;
261
262        let mut hits = Vec::with_capacity(merged.len());
263        for (i, result) in merged.into_iter().enumerate() {
264            let title = match strip_headline_marks(result.title_headline) {
265                Some(title) => title,
266                // No headline (the query didn't match the title itself): fall back to the
267                // page's plain title rather than showing the model an empty string.
268                None => fallback_titles.get(&result.id).cloned().unwrap_or_default(),
269            };
270            hits.push(SearchHit {
271                page_id: result.id,
272                title,
273                chapter_name: result.chapter_name,
274                url_path: result.url_path,
275                rank: result.rank,
276                snippet: strip_headline_marks(result.content_headline),
277                citation_number: starting_number + 1 + i as i32,
278            });
279        }
280
281        Ok(CourseMaterialSearchTool {
282            state: CourseMaterialSearchState {
283                course_id,
284                course_name: course.name,
285                course_slug: course.slug,
286                hits,
287                document_url_prefix: build_courses_base_url(
288                    &app_config.base_url,
289                    &organization.slug,
290                ),
291            },
292        })
293    }
294
295    fn output(&self) -> String {
296        let course = OutputCourse {
297            id: self.state.course_id,
298            name: &self.state.course_name,
299            slug: &self.state.course_slug,
300        };
301        let results: Vec<OutputResult> = self
302            .state
303            .hits
304            .iter()
305            .map(|hit| OutputResult {
306                page_id: hit.page_id,
307                title: &hit.title,
308                chapter_name: hit.chapter_name.as_deref(),
309                url_path: &hit.url_path,
310                rank: hit.rank,
311                snippet: hit.snippet.as_deref(),
312                citation_number: hit.citation_number,
313            })
314            .collect();
315
316        let note = results.is_empty().then_some(
317            "No page matched. Try different wording, or list the course's pages with course_structure.",
318        );
319        let output = Output {
320            course,
321            results,
322            note,
323        };
324        serde_json::to_string_pretty(&output).unwrap_or_else(|_| "No results.".to_string())
325    }
326
327    fn output_description_instructions(&self) -> Option<String> {
328        Some("Quote the material verbatim and name the page title. Cite a page by writing 【0:N†source】 immediately after the sentence that uses it, where N is that result's citation_number; the admin sees those as clickable links to the page. Fetch the whole page with document_lookup (course_id plus the result's page_id) only when the snippet was not clearly enough. If nothing matched, say so plainly instead of guessing.".to_string())
329    }
330
331    /// Column widths in `chatbot_conversation_messages_citations` are `VARCHAR(255)`; truncate to
332    /// fit the way `to_chatbot_conversation_message_citation` does for the Azure path.
333    fn citations(&self) -> Vec<ToolCitation> {
334        self.state
335            .hits
336            .iter()
337            .map(|hit| {
338                let title = truncate_utf8_at_boundary(&hit.title, 255).to_string();
339                let snippet = hit
340                    .snippet
341                    .as_deref()
342                    .map(|s| truncate_utf8_at_boundary(s, 255).to_string())
343                    .unwrap_or_default();
344                let document_url = truncate_utf8_at_boundary(
345                    &format!("{}{}", self.state.document_url_prefix, hit.url_path),
346                    255,
347                )
348                .to_string();
349                ToolCitation {
350                    page_id: hit.page_id,
351                    title,
352                    snippet,
353                    document_url,
354                    citation_number: hit.citation_number,
355                }
356            })
357            .collect()
358    }
359}