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}