Skip to main content

headless_lms_chatbot/chatbot_tools/custom_tools/
find_course.rs

1use 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            // Any candidate id in the group works: the language-versions page lists the whole sibling set.
345            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}