headless_lms_chatbot/chatbot_tools/custom_tools/
course_finder.rs1use headless_lms_utils::cache::Cache;
2use std::collections::HashMap;
3
4use indexmap::IndexMap;
5use serde::Deserializer;
6
7use crate::{
8 azure_chatbot::azure::tools::{AzureLLMFunctionToolDefinition, LLMToolType},
9 chatbot_tools::{
10 ChatbotTool, ChatbotToolDeclaration, ToolProperties, tool_authorization::ToolRequirement,
11 },
12 prelude::*,
13 user_context::ChatbotTurnContext,
14};
15use headless_lms_models::{
16 chatbot_configurations::ToolCategory,
17 organizations::{self, DatabaseOrganization},
18};
19use headless_lms_models::{
20 course_audiences::get_course_ids_by_audience_vectors,
21 course_prerequisites::get_course_ids_by_prerequisite_vectors,
22 courses::{self, Course, get_by_description_vectors},
23 external_courses::{ExternalCourseOutput, get_external_courses_by_embeddings},
24};
25use headless_lms_utils::{
26 azure_embedding::create_embeddings,
27 course_url::build_course_url,
28 json_schema_types::{Schema, string_array_property},
29};
30
31#[derive(Debug, Serialize)]
32pub struct CourseFinderState {
33 courses: Vec<CourseOccurrences>,
34 external_courses: Vec<ExternalCourseOutput>,
35}
36
37#[derive(Deserialize, Clone, Debug)]
38pub struct CourseFinderArguments {
39 #[serde(deserialize_with = "empty_vec_as_none")]
40 description: Option<Vec<String>>,
41 #[serde(deserialize_with = "empty_vec_as_none")]
42 prerequisites: Option<Vec<String>>,
43 #[serde(deserialize_with = "empty_vec_as_none")]
44 audiences: Option<Vec<String>>,
45}
46#[derive(Serialize, Deserialize, Clone, Debug)]
47pub struct CourseOccurrences {
48 course: Course,
49 occurrences: usize,
50 course_url: String,
51}
52
53pub type CourseFinderTool = ToolProperties<CourseFinderState>;
54
55impl ChatbotTool for CourseFinderTool {
56 type Arguments = CourseFinderArguments;
57
58 fn call_requirements(
59 _arguments: &Self::Arguments,
60 _user_context: &ChatbotTurnContext,
61 ) -> Vec<ToolRequirement> {
62 Vec::new()
63 }
64
65 async fn from_db_and_arguments(
66 conn: &mut PgConnection,
67 app_config: &ApplicationConfiguration,
68 _cache: &Cache,
69 arguments: Self::Arguments,
70 _user_context: &ChatbotTurnContext,
71 ) -> ChatbotResult<Self> {
72 let audience_courses = if let Some(audiences) = &arguments.audiences {
73 let audience_embeddings = create_embeddings(app_config, audiences.clone())
74 .await?
75 .to_owned();
76
77 get_course_ids_by_audience_vectors(conn, audience_embeddings, audiences.clone()).await?
78 } else {
79 vec![]
80 };
81
82 let prerequisite_courses = if let Some(prerequisites) = &arguments.prerequisites {
83 let prerequisite_embeddings = create_embeddings(app_config, prerequisites.clone())
84 .await?
85 .to_owned();
86
87 get_course_ids_by_prerequisite_vectors(
88 conn,
89 prerequisite_embeddings,
90 prerequisites.clone(),
91 )
92 .await?
93 } else {
94 vec![]
95 };
96
97 let (description_courses, external_courses) =
98 if let Some(description) = &arguments.description {
99 let description_embeddings = create_embeddings(app_config, description.clone())
100 .await?
101 .to_owned();
102
103 let external_courses = get_external_courses_by_embeddings(
104 conn,
105 description.clone(),
106 description_embeddings.clone(),
107 )
108 .await?;
109
110 let courses =
111 get_by_description_vectors(conn, description_embeddings, description.clone())
112 .await?;
113 (courses, external_courses)
114 } else {
115 (vec![], vec![])
116 };
117
118 let course_ids = [description_courses, audience_courses, prerequisite_courses].concat();
119
120 let mut counts: HashMap<Uuid, usize> = HashMap::new();
121
122 for id in &course_ids {
123 *counts.entry(*id).or_insert(0) += 1;
124 }
125
126 let courses = courses::get_by_ids(conn, &course_ids).await?;
127
128 let organization_ids: Vec<Uuid> = courses
129 .iter()
130 .map(|course| course.organization_id)
131 .collect();
132
133 let organizations = organizations::get_by_ids(conn, &organization_ids).await?;
134
135 let organization_by_id: HashMap<Uuid, &DatabaseOrganization> = organizations
136 .iter()
137 .map(|organization| (organization.id, organization))
138 .collect();
139
140 let mut course_occurrences: Vec<CourseOccurrences> = courses
141 .into_iter()
142 .filter_map(|course| {
143 if course.is_draft || course.is_test_mode || course.is_unlisted {
144 return None;
145 }
146 let organization = organization_by_id.get(&course.organization_id)?;
147
148 Some(CourseOccurrences {
149 occurrences: counts[&course.id],
150 course_url: build_course_url(
151 &app_config.base_url,
152 &organization.slug,
153 &course.slug,
154 ),
155 course,
156 })
157 })
158 .collect();
159
160 course_occurrences.sort_by_key(|b| std::cmp::Reverse(b.occurrences));
161
162 Ok(CourseFinderTool {
163 state: CourseFinderState {
164 courses: course_occurrences,
165 external_courses,
166 },
167 })
168 }
169
170 fn output(&self) -> String {
171 serde_json::to_string(&self.state).unwrap_or_else(|_| "No courses found".to_string())
172 }
173
174 fn output_description_instructions(&self) -> Option<String> {
175 Some("Do not return the whole JSON of the courses to the user. Courses under the 'courses'key are what you should prioritize. If in addition to those courses there is some external course that matches the user query especially well, you can recommend that as well. If you recommend an external course, state that an external course is not in courses.mooc.fi, but it might be on an older version of the platform, don't advertise too much that on what platform the course is.. Present the most suitable courses based on the user query. Use the course names and course descriptions to give a list and a very brief and summarized description of each course to the user. If there are duplicate courses ignore them. You can also mention why the course could be suitable to the user based on their request.".to_string())
176 }
177}
178
179impl ChatbotToolDeclaration for CourseFinderTool {
180 const NAME: &'static str = "course_finder";
181
182 fn offer_requirements(_user_context: &ChatbotTurnContext) -> Vec<ToolRequirement> {
183 Vec::new()
184 }
185
186 const CATEGORY: ToolCategory = ToolCategory::CourseCatalog;
187
188 fn get_tool_definition() -> AzureLLMFunctionToolDefinition {
189 AzureLLMFunctionToolDefinition {
190 tool_type: LLMToolType::Function,
191 name: Self::NAME.to_string(),
192 description: "Find suitable courses for the user if they want to find available courses for their conditions. The arguments should be created based on the terms with which the user wants to filter the courses. The needed arguments should therefore be parsed from the user message. The arguments are arrays of keywords for the parameters the user is using to search the courses. At least one of the three arguments is required. Match on any single argument will find a course so it is safe to provide all types of arguments when suitable. This tool is useful to find any courses if the user wants recommendations for courses they can take.".to_string(),
193 parameters: Schema::strict_object(
194 IndexMap::from([
195 (
196 "description".to_string(),
197 string_array_property(Some("List of keywords used to search course descriptions based on if the user tries to find courses based on what they contain or teach.")),
198 ),
199 (
200 "prerequisites".to_string(),
201 string_array_property(Some("List of keywords of preliminary knowledge possessed to be suitable for a course.")),
202 ),
203 (
204 "audiences".to_string(),
205 string_array_property(Some("List of keywords of audience types that a course is suitable for.")),
206 ),
207 ]),
208 None,
209 ),
210 strict: true,
211 }
212 }
213}
214
215fn empty_vec_as_none<'de, D>(deserializer: D) -> Result<Option<Vec<String>>, D::Error>
216where
217 D: Deserializer<'de>,
218{
219 let opt = Option::<Vec<String>>::deserialize(deserializer)?;
220
221 Ok(opt.and_then(|vec| {
222 let vec: Vec<String> = vec
223 .into_iter()
224 .map(|s| s.trim().to_owned())
225 .filter(|s| !s.is_empty())
226 .collect();
227
228 if vec.is_empty() { None } else { Some(vec) }
229 }))
230}