Skip to main content

headless_lms_chatbot/chatbot_tools/custom_tools/
course_finder.rs

1use 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}