headless_lms_chatbot/chatbot_tools/custom_tools/
find_course.rs1use headless_lms_authorization::Action;
2use headless_lms_utils::cache::Cache;
3use std::str::FromStr;
4
5use indexmap::IndexMap;
6
7use headless_lms_models::chatbot_configurations::ToolCategory;
8use headless_lms_models::{
9 course_instances::{self, CourseInstance},
10 courses::{self, Course},
11 organizations,
12};
13use headless_lms_utils::json_schema_types::{JSONType, JsonItem, Schema, SchemaPropertyType};
14
15use crate::{
16 azure_chatbot::azure::tools::{AzureLLMFunctionToolDefinition, LLMToolType},
17 chatbot_tools::{
18 ChatbotTool, ChatbotToolDeclaration, ToolProperties, tool_authorization::ToolRequirement,
19 },
20 prelude::*,
21 user_context::ChatbotTurnContext,
22};
23
24pub type FindCourseTool = ToolProperties<FindCourseState>;
25
26pub struct FindCourseState {
27 candidates: Vec<CourseCandidate>,
28 base_url: String,
29}
30
31struct CourseCandidate {
32 course: Course,
33 instances: Vec<CourseInstance>,
34 organization_name: String,
35}
36
37#[derive(Serialize)]
38struct CourseCandidateOutput {
39 course_id: Uuid,
40 name: String,
41 slug: String,
42 language_code: String,
43 organization_name: String,
44 is_draft: bool,
45 is_test_mode: bool,
46 #[serde(skip_serializing_if = "Option::is_none")]
47 closed_at: Option<chrono::DateTime<chrono::Utc>>,
48 #[serde(skip_serializing_if = "Option::is_none")]
49 closed_additional_message: Option<String>,
50 #[serde(skip_serializing_if = "Option::is_none")]
51 closed_course_successor_id: Option<Uuid>,
52 instances: Vec<CourseInstanceOutput>,
53}
54
55#[derive(Serialize)]
56struct CourseInstanceOutput {
57 course_instance_id: Uuid,
58 #[serde(skip_serializing_if = "Option::is_none")]
59 name: Option<String>,
60 #[serde(skip_serializing_if = "Option::is_none")]
61 starts_at: Option<chrono::DateTime<chrono::Utc>>,
62 #[serde(skip_serializing_if = "Option::is_none")]
63 ends_at: Option<chrono::DateTime<chrono::Utc>>,
64 #[serde(skip_serializing_if = "Option::is_none")]
65 support_email: Option<String>,
66}
67
68impl From<&CourseInstance> for CourseInstanceOutput {
69 fn from(instance: &CourseInstance) -> Self {
70 Self {
71 course_instance_id: instance.id,
72 name: instance.name.clone(),
73 starts_at: instance.starts_at,
74 ends_at: instance.ends_at,
75 support_email: instance.support_email.clone(),
76 }
77 }
78}
79
80impl From<&CourseCandidate> for CourseCandidateOutput {
81 fn from(candidate: &CourseCandidate) -> Self {
82 let course = &candidate.course;
83 Self {
84 course_id: course.id,
85 name: course.name.clone(),
86 slug: course.slug.clone(),
87 language_code: course.language_code.clone(),
88 organization_name: candidate.organization_name.clone(),
89 is_draft: course.is_draft,
90 is_test_mode: course.is_test_mode,
91 closed_at: course.closed_at,
92 closed_additional_message: course.closed_additional_message.clone(),
93 closed_course_successor_id: course.closed_course_successor_id,
94 instances: candidate
95 .instances
96 .iter()
97 .map(CourseInstanceOutput::from)
98 .collect(),
99 }
100 }
101}
102
103#[derive(Deserialize)]
104pub struct FindCourseArguments {
105 query: String,
106}
107
108const MAX_CANDIDATES: i64 = 5;
109
110impl ChatbotToolDeclaration for FindCourseTool {
111 const NAME: &'static str = "find_course";
112
113 fn offer_requirements(_user_context: &ChatbotTurnContext) -> Vec<ToolRequirement> {
114 vec![ToolRequirement::global(Action::Administrate)]
115 }
116
117 const CATEGORY: ToolCategory = ToolCategory::AdminSupportCourses;
118
119 fn get_tool_definition() -> AzureLLMFunctionToolDefinition {
120 AzureLLMFunctionToolDefinition {
121 tool_type: LLMToolType::Function,
122 name: Self::NAME.to_string(),
123 description: "Find a course by UUID, exact slug, or (part of) its name. Use this to resolve which course an admin means before calling course- or user-scoped tools.".to_string(),
124 parameters: Schema::strict_object(
125 IndexMap::from([(
126 "query".to_string(),
127 SchemaPropertyType::Item(JsonItem {
128 type_field: JSONType::String,
129 description: Some(
130 "Course UUID, exact slug, or (part of) the course name.".to_string(),
131 ),
132 }),
133 )]),
134 None,
135 ),
136 strict: true,
137 }
138 }
139}
140
141impl ChatbotTool for FindCourseTool {
142 type Arguments = FindCourseArguments;
143
144 fn call_requirements(
145 _arguments: &Self::Arguments,
146 _user_context: &ChatbotTurnContext,
147 ) -> Vec<ToolRequirement> {
148 vec![ToolRequirement::global(Action::Administrate)]
149 }
150
151 fn parse_arguments(args_string: String) -> ChatbotResult<Self::Arguments> {
152 let mut arguments: Self::Arguments = serde_json::from_str(&args_string).map_err(|e| {
153 chatbot_err!(
154 InvalidToolArguments,
155 format!("Couldn't parse tool arguments. Arguments: {args_string}"),
156 e
157 )
158 })?;
159 arguments.query = arguments.query.trim().to_string();
160 if arguments.query.is_empty() {
161 return Err(chatbot_err!(
162 InvalidToolArguments,
163 "query must not be empty.".to_string()
164 ));
165 }
166 Ok(arguments)
167 }
168
169 async fn from_db_and_arguments(
170 conn: &mut PgConnection,
171 app_config: &ApplicationConfiguration,
172 _cache: &Cache,
173 arguments: Self::Arguments,
174 _user_context: &ChatbotTurnContext,
175 ) -> ChatbotResult<Self> {
176 let base_url = app_config.base_url.trim_end_matches('/').to_string();
177 let query = arguments.query;
178
179 let courses = if let Ok(course_id) = Uuid::from_str(&query) {
180 match courses::get_course(conn, course_id).await.optional()? {
181 Some(course) => vec![course],
182 None => {
183 courses::search_courses_by_slug_or_name(conn, &query, MAX_CANDIDATES).await?
184 }
185 }
186 } else {
187 courses::search_courses_by_slug_or_name(conn, &query, MAX_CANDIDATES).await?
188 };
189
190 let organization_ids: Vec<Uuid> = courses.iter().map(|c| c.organization_id).collect();
191 let organization_names: std::collections::HashMap<Uuid, String> =
192 organizations::get_by_ids(conn, &organization_ids)
193 .await?
194 .into_iter()
195 .map(|org| (org.id, org.name))
196 .collect();
197
198 let mut candidates = Vec::with_capacity(courses.len());
199 for course in courses {
200 let instances =
201 course_instances::get_course_instances_for_course(conn, course.id).await?;
202 let organization_name = organization_names
203 .get(&course.organization_id)
204 .cloned()
205 .unwrap_or_default();
206 candidates.push(CourseCandidate {
207 course,
208 instances,
209 organization_name,
210 });
211 }
212
213 Ok(FindCourseTool {
214 state: FindCourseState {
215 candidates,
216 base_url,
217 },
218 })
219 }
220
221 fn output(&self) -> String {
222 let candidates: Vec<CourseCandidateOutput> = self
223 .state
224 .candidates
225 .iter()
226 .map(CourseCandidateOutput::from)
227 .collect();
228 serde_json::to_string_pretty(&candidates).unwrap_or_else(|_| "No courses found".to_string())
229 }
230
231 fn output_description_instructions(&self) -> Option<String> {
232 let candidates = &self.state.candidates;
233 let base_url = &self.state.base_url;
234 let mut notes = vec![
235 "If several courses match (e.g. language versions of the same course — compare \
236 language_code), ask the admin which one before proceeding. The instance contact \
237 emails shown here may be stale; the course_configuration tool's staff facet is the \
238 fresher source."
239 .to_string(),
240 ];
241
242 if !candidates.is_empty() {
243 let overview_links = candidates
244 .iter()
245 .map(|c| {
246 format!(
247 "{} ({}): {base_url}/manage/courses/{}/overview",
248 c.course.name, c.course.language_code, c.course.id
249 )
250 })
251 .collect::<Vec<_>>()
252 .join(", ");
253 notes.push(format!(
254 "Course overview pages, to confirm you and the admin are looking at the same \
255 course: {overview_links}."
256 ));
257 }
258
259 if candidates.is_empty() {
260 notes.push(
261 "No courses matched. This means either nothing matched the query, or the \
262 course was deleted (deleted courses are excluded from this search)."
263 .to_string(),
264 );
265 }
266
267 if candidates.len() == MAX_CANDIDATES as usize {
268 notes.push(format!(
269 "Results are capped at {MAX_CANDIDATES} and ordered exact slug match > name \
270 substring > fuzzy match; there may be more matching courses that were \
271 silently truncated from this list."
272 ));
273 }
274
275 if candidates
276 .iter()
277 .any(|c| c.course.is_test_mode || c.course.is_draft)
278 {
279 notes.push(
280 "Some results have is_test_mode or is_draft set. A test-mode course is a \
281 staff testing copy and a draft course is unpublished — neither is the course \
282 a student is asking about."
283 .to_string(),
284 );
285 }
286
287 if candidates.iter().any(|c| c.course.closed_at.is_some()) {
288 let now = chrono::Utc::now();
289 let mut closed_at_note = String::from(
290 "closed_at is a scheduled closing timestamp: absent means the course was \
291 never scheduled to close, a future value means it's still open, and only a \
292 past value means it's actually closed.",
293 );
294 if candidates.iter().any(|c| {
295 c.course.closed_at.is_some_and(|t| t <= now)
296 && c.course.closed_course_successor_id.is_none()
297 }) {
298 closed_at_note.push_str(
299 " A closed course with no closed_course_successor_id has nowhere \
300 configured to send the student.",
301 );
302 }
303 if candidates
304 .iter()
305 .any(|c| c.course.closed_course_successor_id.is_some())
306 {
307 closed_at_note.push_str(
308 " closed_course_successor_id is a course id, not a name — call \
309 find_course again to identify it.",
310 );
311 }
312 notes.push(closed_at_note);
313 }
314
315 if candidates
316 .iter()
317 .flat_map(|c| &c.instances)
318 .any(|i| i.starts_at.is_none() || i.ends_at.is_none())
319 {
320 notes.push(
321 "Some instances are missing starts_at or ends_at. The platform itself is \
322 inconsistent about whether such an instance counts as started, so report the \
323 absence rather than asserting whether the instance is running."
324 .to_string(),
325 );
326 }
327
328 let mut name_to_candidates: std::collections::HashMap<&str, Vec<&CourseCandidate>> =
329 std::collections::HashMap::new();
330 for candidate in candidates {
331 name_to_candidates
332 .entry(candidate.course.name.as_str())
333 .or_default()
334 .push(candidate);
335 }
336 if let Some(ambiguous) = name_to_candidates.values().find(|group| {
337 group
338 .iter()
339 .map(|c| c.course.language_code.as_str())
340 .collect::<std::collections::HashSet<_>>()
341 .len()
342 >= 2
343 }) {
344 let representative_id = ambiguous[0].course.id;
346 notes.push(format!(
347 "Some results share a name but differ in language_code — these are separate \
348 course rows, and a student's progress lives in exactly one of them. \
349 {base_url}/manage/courses/{representative_id}/language-versions lists the \
350 whole sibling set side by side."
351 ));
352 }
353
354 Some(notes.join(" "))
355 }
356}