Skip to main content

headless_lms_models/
course_prerequisites.rs

1use crate::prelude::*;
2use pgvector::Vector;
3use serde::{Deserialize, Serialize};
4use utoipa::ToSchema;
5
6#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, ToSchema)]
7pub struct CoursePrerequisite {
8    pub id: Uuid,
9    pub created_at: DateTime<Utc>,
10    pub updated_at: DateTime<Utc>,
11    pub deleted_at: Option<DateTime<Utc>>,
12    pub course_id: Uuid,
13    pub prerequisite: String,
14    #[schema(value_type = Option<Vec<f32>>)]
15    pub embedding: Option<Vector>,
16}
17
18#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize, ToSchema, Hash)]
19pub struct EditCoursePrerequisite {
20    pub id: Uuid,
21    pub course_id: Uuid,
22    pub prerequisite: String,
23}
24
25pub async fn insert_course_prerequisites(
26    conn: &mut PgConnection,
27    course_id: Uuid,
28    new_prerequisites: Vec<String>,
29    embeddings: Vec<Vec<f32>>,
30) -> ModelResult<Vec<CoursePrerequisite>> {
31    let embed_vecs: Vec<Vector> = embeddings.into_iter().map(Vector::from).collect();
32
33    let res = sqlx::query_as!(
34        CoursePrerequisite,
35        r#"
36INSERT INTO course_prerequisites (
37    course_id,
38    prerequisite,
39    embedding
40  )
41SELECT $1,
42       t.prerequisite,
43       t.embedding
44FROM UNNEST(
45    $2::text[],
46    $3::vector[]
47) AS t(prerequisite, embedding)
48RETURNING *
49    "#,
50        course_id,
51        &new_prerequisites,
52        &embed_vecs as _
53    )
54    .fetch_all(conn)
55    .await?;
56    Ok(res)
57}
58
59pub async fn get_all_edit_course_prerequisites(
60    conn: &mut PgConnection,
61) -> ModelResult<Vec<EditCoursePrerequisite>> {
62    let res = sqlx::query_as!(
63        EditCoursePrerequisite,
64        "
65SELECT id,
66prerequisite,
67course_id
68FROM course_prerequisites
69WHERE deleted_at IS NULL
70",
71    )
72    .fetch_all(conn)
73    .await?;
74    Ok(res)
75}
76
77pub async fn get_edit_course_prerequisites_by_course_id(
78    conn: &mut PgConnection,
79    course_id: Uuid,
80) -> ModelResult<Vec<EditCoursePrerequisite>> {
81    let res = sqlx::query_as!(
82        EditCoursePrerequisite,
83        "
84SELECT id,
85prerequisite,
86course_id
87FROM course_prerequisites
88WHERE course_id = $1
89AND deleted_at IS NULL
90",
91        course_id
92    )
93    .fetch_all(conn)
94    .await?;
95    Ok(res)
96}
97
98pub async fn get_by_course_id(
99    conn: &mut PgConnection,
100    course_id: Uuid,
101) -> ModelResult<Vec<CoursePrerequisite>> {
102    let res = sqlx::query_as!(
103        CoursePrerequisite,
104        r#"
105SELECT *
106FROM course_prerequisites
107WHERE course_id = $1
108AND deleted_at IS NULL
109"#,
110        course_id
111    )
112    .fetch_all(conn)
113    .await?;
114    Ok(res)
115}
116
117pub async fn upsert_course_prerequisites(
118    conn: &mut PgConnection,
119    course_id: Uuid,
120    prerequisite_ids: Vec<Uuid>,
121    updated_prerequisites: Vec<String>,
122    embeddings: Vec<Vec<f32>>,
123) -> ModelResult<Vec<CoursePrerequisite>> {
124    let embed_vecs: Vec<Vector> = embeddings.into_iter().map(Vector::from).collect();
125
126    let id_count = sqlx::query_scalar!(
127        r#"
128SELECT COUNT(*) AS "count!"
129FROM course_prerequisites
130WHERE id = ANY($2)
131  AND course_id != $1
132"#,
133        course_id,
134        &prerequisite_ids,
135    )
136    .fetch_one(&mut *conn)
137    .await?;
138
139    if id_count != 0 {
140        return Err(model_err!(
141            InvalidRequest,
142            "Ids of some given prerequisite entries already exists on other courses.".to_string()
143        ));
144    }
145
146    let res = sqlx::query_as!(
147        CoursePrerequisite,
148        r#"
149INSERT INTO course_prerequisites (course_id, id, prerequisite, embedding)
150SELECT $1,
151  course_prerequisite.id,
152  course_prerequisite.prerequisite,
153  course_prerequisite.embedding
154FROM UNNEST ($2::UUID [], $3::TEXT [], $4::VECTOR []) AS course_prerequisite(id, prerequisite, embedding) ON CONFLICT (id) DO
155UPDATE
156SET prerequisite = EXCLUDED.prerequisite,
157  embedding = EXCLUDED.embedding
158WHERE course_prerequisites.deleted_at IS NULL
159RETURNING *
160"#,
161        course_id,
162        &prerequisite_ids,
163        &updated_prerequisites,
164        &embed_vecs as _
165    )
166    .fetch_all(conn)
167    .await?;
168
169    Ok(res)
170}
171
172pub async fn delete_batch(
173    conn: &mut PgConnection,
174    ids_to_delete: Vec<Uuid>,
175) -> ModelResult<Vec<CoursePrerequisite>> {
176    let res = sqlx::query_as!(
177        CoursePrerequisite,
178        r#"
179UPDATE course_prerequisites
180SET deleted_at = now()
181WHERE id = ANY($1::UUID [])
182AND deleted_at IS NULL
183RETURNING *
184"#,
185        &ids_to_delete
186    )
187    .fetch_all(conn)
188    .await?;
189    Ok(res)
190}
191
192pub async fn get_course_ids_by_prerequisite_vectors(
193    conn: &mut PgConnection,
194    prerequisite_vecs: Vec<Vec<f32>>,
195    prerequisite_keywords: Vec<String>,
196) -> ModelResult<Vec<Uuid>> {
197    let vectors: Vec<Vector> = prerequisite_vecs.into_iter().map(Vector::from).collect();
198    let res = sqlx::query_scalar!(
199        r#"
200SELECT course_id
201FROM (
202    SELECT
203        p.course_id,
204        MIN(p.embedding <#> v.embedding) AS distance
205    FROM course_prerequisites p
206    CROSS JOIN unnest($1::vector[]) AS v(embedding)
207    WHERE deleted_at IS NULL
208    GROUP BY p.course_id
209    ORDER BY distance ASC
210    LIMIT 5
211) t
212UNION ALL
213SELECT DISTINCT p.course_id
214FROM course_prerequisites p
215CROSS JOIN unnest($2::text[]) AS k(keyword)
216WHERE deleted_at IS NULL
217AND p.prerequisite % k.keyword
218        "#,
219        &vectors as _,
220        &prerequisite_keywords
221    )
222    .fetch_all(conn)
223    .await?;
224    Ok(res.into_iter().flatten().collect())
225}