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