Skip to main content

headless_lms_models/
course_audiences.rs

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}