headless_lms_chatbot/chatbot_tools/custom_tools/
course_material_search.rs1use 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
88impl<'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
131fn 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 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 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 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}