diff --git a/parakeet/src/loaders/labeler.rs b/parakeet/src/loaders/labeler.rs index e520ecdf..653fc824 100644 --- a/parakeet/src/loaders/labeler.rs +++ b/parakeet/src/loaders/labeler.rs @@ -19,18 +19,15 @@ pub fn build_labeler_records_query(actor_ids_str: &str) -> String { ) } -/// Build SQL query for loading labels by URI (DENORMALIZED) +/// Build SQL query for loading labels by actor_id (DENORMALIZED) /// /// Labels are now stored as actor_label[] arrays on the actors table. /// Each actor has a labels array containing labels they've applied. +/// Note: URI parsing done in Rust, query by actor_id for efficiency /// /// This function is public for testing purposes. pub fn build_labels_query() -> &'static str { - "WITH target_actor AS ( - SELECT id FROM actors - WHERE did = SPLIT_PART(SUBSTRING($1 FROM 6), '/', 1) - ) - SELECT + "SELECT (lbl).labeler_actor_id, (lbl).label as label, $1::text as uri, @@ -39,30 +36,25 @@ pub fn build_labels_query() -> &'static str { (lbl).negated, (lbl).expires, NULL::bytea as sig, - (lbl).created_at, - labeler.did as labeler - FROM target_actor ta - CROSS JOIN actors labeler_subjects + (lbl).created_at + FROM actors labeler_subjects CROSS JOIN unnest(labeler_subjects.labels) AS lbl - INNER JOIN actors labeler ON (lbl).labeler_actor_id = labeler.id - WHERE labeler_subjects.id = ta.id + WHERE labeler_subjects.id = $2 AND (lbl).negated = false - AND labeler.did = ANY($2) + AND (lbl).labeler_actor_id = ANY($3) ORDER BY (lbl).created_at" } -/// Build SQL query for batch loading labels by multiple URIs (DENORMALIZED) +/// Build SQL query for batch loading labels by actor_ids (DENORMALIZED) /// /// Labels are now stored as actor_label[] arrays on the actors table. -/// This query unnests labels from multiple target actors. +/// Note: URI parsing done in Rust, query by actor_ids for efficiency /// /// This function is public for testing purposes. pub fn build_labels_many_query() -> &'static str { "WITH target_actors AS ( SELECT unnest($1::text[]) as uri, - a.id as actor_id - FROM unnest($1::text[]) uri_val - INNER JOIN actors a ON a.did = SPLIT_PART(SUBSTRING(uri_val FROM 6), '/', 1) + unnest($2::int[]) as actor_id ) SELECT (lbl).labeler_actor_id, @@ -73,14 +65,12 @@ pub fn build_labels_many_query() -> &'static str { (lbl).negated, (lbl).expires, NULL::bytea as sig, - (lbl).created_at, - labeler.did as labeler + (lbl).created_at FROM target_actors ta INNER JOIN actors labeler_subjects ON labeler_subjects.id = ta.actor_id CROSS JOIN unnest(labeler_subjects.labels) AS lbl - INNER JOIN actors labeler ON (lbl).labeler_actor_id = labeler.id WHERE (lbl).negated = false - AND labeler.did = ANY($2) + AND (lbl).labeler_actor_id = ANY($3) ORDER BY (lbl).created_at" } @@ -99,16 +89,29 @@ pub struct EnrichedLabeler { pub like_count: i32, } -pub struct LabelServiceLoader(pub(super) Pool); +pub struct LabelServiceLoader( + pub(super) Pool, + pub(super) std::sync::Arc, +); pub type LabelServiceLoaderRet = (EnrichedLabeler, Vec); impl BatchFn for LabelServiceLoader { async fn load(&mut self, keys: &[String]) -> HashMap { let mut conn = self.0.get().await.unwrap(); - // Load labelers from actors table (actors with labeler_cid IS NOT NULL) + // Resolve DIDs to actor_ids using IdCache + let did_to_actor = self.1.get_actor_ids(keys).await; + + // Collect actor_ids for query + let actor_ids: Vec = did_to_actor.values().map(|a| a.actor_id).collect(); + + if actor_ids.is_empty() { + return HashMap::new(); + } + + // Load labelers from actors table by actor_id (more efficient than DID) let actors: Vec = diesel_async::RunQueryDsl::load( schema::actors::table - .filter(schema::actors::did.eq_any(keys)) + .filter(schema::actors::id.eq_any(&actor_ids)) .filter(schema::actors::labeler_cid.is_not_null()) .filter(schema::actors::labeler_status.eq(parakeet_db::types::LabelerStatus::Complete)) .select(models::Actor::as_select()), @@ -167,6 +170,32 @@ impl LabelLoader { pub async fn load(&self, uri: &str, services: &[LabelConfigItem]) -> Vec { let mut conn = self.0.get().await.unwrap(); + // Parse URI to extract DID (at://did:plc:xxx/...) + let subject_did = if uri.starts_with("at://") { + uri.strip_prefix("at://") + .and_then(|s| s.split('/').next()) + } else { + None + }; + + let subject_did = match subject_did { + Some(did) => did, + None => { + tracing::warn!("Invalid AT URI format for labels: {}", uri); + return Vec::new(); + } + }; + + // Resolve subject DID to actor_id via IdCache + let subject_actor_id = match self.1.get_actor_id_only(subject_did).await { + Some(actor_id) => actor_id, + None => { + tracing::debug!("Subject actor not found for labels: {}", subject_did); + return Vec::new(); + } + }; + + // Resolve service DIDs to actor_ids let service_dids: Vec<&str> = services .iter() .map(|v| v.labeler.as_str()) @@ -176,6 +205,13 @@ impl LabelLoader { return Vec::new(); } + let service_actors = self.1.get_actor_ids(&service_dids.iter().map(|s| s.to_string()).collect::>()).await; + let labeler_actor_ids: Vec = service_actors.values().map(|a| a.actor_id).collect(); + + if labeler_actor_ids.is_empty() { + return Vec::new(); + } + let query = build_labels_query(); #[derive(diesel::QueryableByName)] @@ -198,14 +234,13 @@ impl LabelLoader { sig: Option>, #[diesel(sql_type = diesel::sql_types::Timestamptz)] created_at: chrono::DateTime, - #[diesel(sql_type = diesel::sql_types::Text)] - labeler: String, } let labels: Vec = diesel_async::RunQueryDsl::load( diesel::sql_query(query) .bind::(uri) - .bind::, _>(&service_dids), + .bind::(subject_actor_id) + .bind::, _>(&labeler_actor_ids), &mut conn, ) .await @@ -214,19 +249,36 @@ impl LabelLoader { vec![] }); + // Resolve labeler_actor_ids back to DIDs for the result + let unique_labeler_ids: Vec = labels + .iter() + .map(|row| row.labeler_actor_id) + .collect::>() + .into_iter() + .collect(); + + let labeler_data = self.1.get_actor_data_many(&unique_labeler_ids).await; + labels .into_iter() - .map(|row| models::Label { - labeler_actor_id: row.labeler_actor_id, - label: row.label, - uri: row.uri, - self_label: row.self_label, - cid: row.cid, - negated: row.negated, - expires: row.expires, - sig: row.sig, - created_at: row.created_at, - labeler: row.labeler, + .map(|row| { + let labeler_did = labeler_data + .get(&row.labeler_actor_id) + .map(|d| d.did.clone()) + .unwrap_or_else(|| String::from("did:unknown")); + + models::Label { + labeler_actor_id: row.labeler_actor_id, + label: row.label, + uri: row.uri, + self_label: row.self_label, + cid: row.cid, + negated: row.negated, + expires: row.expires, + sig: row.sig, + created_at: row.created_at, + labeler: labeler_did, + } }) .collect() } @@ -238,17 +290,55 @@ impl LabelLoader { ) -> HashMap> { let mut conn = self.0.get().await.unwrap(); - let service_dids: Vec<&str> = services + if services.is_empty() || uris.is_empty() { + return HashMap::new(); + } + + // Parse URIs to extract DIDs and build parallel arrays + let mut valid_uris = Vec::new(); + let mut subject_dids = Vec::new(); + for uri in uris { + if let Some(did) = uri.strip_prefix("at://").and_then(|s| s.split('/').next()) { + valid_uris.push(uri.as_str()); + subject_dids.push(did.to_string()); + } else { + tracing::warn!("Invalid AT URI format for labels: {}", uri); + } + } + + if valid_uris.is_empty() { + return HashMap::new(); + } + + // Resolve subject DIDs to actor_ids + let subject_actors = self.1.get_actor_ids(&subject_dids).await; + let mut uri_actor_pairs: Vec<(&str, i32)> = Vec::new(); + for (uri, did) in valid_uris.iter().zip(subject_dids.iter()) { + if let Some(actor) = subject_actors.get(did) { + uri_actor_pairs.push((uri, actor.actor_id)); + } + } + + if uri_actor_pairs.is_empty() { + return HashMap::new(); + } + + let uri_refs: Vec<&str> = uri_actor_pairs.iter().map(|(uri, _)| *uri).collect(); + let actor_ids: Vec = uri_actor_pairs.iter().map(|(_, id)| *id).collect(); + + // Resolve service DIDs to actor_ids + let service_dids: Vec = services .iter() - .map(|v| v.labeler.as_str()) + .map(|v| v.labeler.clone()) .collect(); - if service_dids.is_empty() || uris.is_empty() { + let service_actors = self.1.get_actor_ids(&service_dids).await; + let labeler_actor_ids: Vec = service_actors.values().map(|a| a.actor_id).collect(); + + if labeler_actor_ids.is_empty() { return HashMap::new(); } - let uri_refs: Vec<&str> = uris.iter().map(|s| s.as_str()).collect(); - let query = build_labels_many_query(); #[derive(diesel::QueryableByName)] @@ -271,14 +361,13 @@ impl LabelLoader { sig: Option>, #[diesel(sql_type = diesel::sql_types::Timestamptz)] created_at: chrono::DateTime, - #[diesel(sql_type = diesel::sql_types::Text)] - labeler: String, } let labels: Vec = diesel_async::RunQueryDsl::load( diesel::sql_query(query) .bind::, _>(&uri_refs) - .bind::, _>(&service_dids), + .bind::, _>(&actor_ids) + .bind::, _>(&labeler_actor_ids), &mut conn, ) .await @@ -287,19 +376,36 @@ impl LabelLoader { vec![] }); + // Resolve labeler_actor_ids back to DIDs + let unique_labeler_ids: Vec = labels + .iter() + .map(|row| row.labeler_actor_id) + .collect::>() + .into_iter() + .collect(); + + let labeler_data = self.1.get_actor_data_many(&unique_labeler_ids).await; + labels .into_iter() - .map(|row| models::Label { - labeler_actor_id: row.labeler_actor_id, - label: row.label.clone(), - uri: row.uri.clone(), - self_label: row.self_label, - cid: row.cid, - negated: row.negated, - expires: row.expires, - sig: row.sig, - created_at: row.created_at, - labeler: row.labeler, + .map(|row| { + let labeler_did = labeler_data + .get(&row.labeler_actor_id) + .map(|d| d.did.clone()) + .unwrap_or_else(|| String::from("did:unknown")); + + models::Label { + labeler_actor_id: row.labeler_actor_id, + label: row.label.clone(), + uri: row.uri.clone(), + self_label: row.self_label, + cid: row.cid, + negated: row.negated, + expires: row.expires, + sig: row.sig, + created_at: row.created_at, + labeler: labeler_did, + } }) .into_group_map_by(|v| v.uri.clone()) } diff --git a/parakeet/src/loaders/mod.rs b/parakeet/src/loaders/mod.rs index 66e7b23d..5b1f2a5b 100644 --- a/parakeet/src/loaders/mod.rs +++ b/parakeet/src/loaders/mod.rs @@ -98,7 +98,7 @@ impl Dataloaders { profile_by_id: new_plc_loader(ProfileByIdLoader(pool.clone(), id_cache.clone()), "profile_id:", 3600, 50_000), // 1 hour TTL, 10k capacity for feeds/lists/etc feedgen: new_plc_loader(FeedGenLoader(pool.clone(), id_cache.clone()), "feedgen:", 3600, 10_000), - labeler: new_plc_loader(LabelServiceLoader(pool.clone()), "labeler:", 3600, 10_000), + labeler: new_plc_loader(LabelServiceLoader(pool.clone(), id_cache.clone()), "labeler:", 3600, 10_000), list: new_plc_loader(ListLoader(pool.clone()), "list:", 3600, 10_000), starterpacks: new_plc_loader(StarterPackLoader(pool.clone(), id_cache.clone()), "starterpacks:", 3600, 10_000), verification: new_plc_loader(VerificationLoader(pool.clone(), id_cache.clone()), "verification:", 3600, 10_000),