diff --git a/Cargo.lock b/Cargo.lock index 85156764..83fdcb20 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2748,6 +2748,7 @@ dependencies = [ "multibase", "parakeet-db", "parakeet-index", + "reqwest", "serde", "serde_json", "tokio", diff --git a/lexica/src/app_bsky/feed.rs b/lexica/src/app_bsky/feed.rs index 487a7c95..45fd4c7f 100644 --- a/lexica/src/app_bsky/feed.rs +++ b/lexica/src/app_bsky/feed.rs @@ -5,7 +5,7 @@ use crate::app_bsky::graph::ListViewBasic; use crate::app_bsky::richtext::FacetMain; use crate::com_atproto::label::Label; use chrono::prelude::*; -use serde::Serialize; +use serde::{Deserialize, Serialize}; use std::str::FromStr; #[derive(Clone, Debug, Serialize)] @@ -39,8 +39,8 @@ pub struct FeedViewPost { pub reply: Option, #[serde(skip_serializing_if = "Option::is_none")] pub reason: Option, - // #[serde(skip_serializing_if = "Option::is_none")] - // pub feed_context: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub feed_context: Option, } #[derive(Debug, Serialize)] @@ -75,15 +75,22 @@ pub enum ReplyRefPost { #[serde(tag = "$type")] pub enum FeedViewPostReason { #[serde(rename = "app.bsky.feed.defs#reasonRepost")] - Repost { - by: ProfileViewBasic, - #[serde(rename = "indexedAt")] - indexed_at: DateTime, - }, + Repost(FeedReasonRepost), #[serde(rename = "app.bsky.feed.defs#reasonPin")] Pin, } +#[derive(Debug, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct FeedReasonRepost { + pub by: ProfileViewBasic, + #[serde(skip_serializing_if = "Option::is_none")] + pub uri: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cid: Option, + pub indexed_at: DateTime, +} + #[derive(Debug, Serialize)] #[serde(rename_all = "camelCase")] pub struct ThreadViewPost { @@ -185,3 +192,34 @@ pub struct Like { pub created_at: DateTime, pub indexed_at: DateTime, } + +#[derive(Clone, Debug, Deserialize, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct FeedSkeletonResponse { + pub feed: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub cursor: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub req_id: Option, +} + +#[derive(Clone, Debug, Deserialize, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct SkeletonFeedPost { + pub post: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub reason: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub feed_context: Option, +} + +#[derive(Clone, Debug, Deserialize, Serialize)] +#[serde(tag = "$type")] +pub enum SkeletonReason { + #[serde(rename = "app.bsky.feed.defs#skeletonReasonPin")] + Pin {}, + #[serde(rename = "app.bsky.feed.defs#skeletonReasonRepost")] + Repost { + repost: String, + }, +} diff --git a/parakeet/Cargo.toml b/parakeet/Cargo.toml index 59b3cef7..3d1a8608 100644 --- a/parakeet/Cargo.toml +++ b/parakeet/Cargo.toml @@ -23,6 +23,7 @@ lexica = { path = "../lexica" } multibase = "0.9.1" parakeet-db = { path = "../parakeet-db" } parakeet-index = { path = "../parakeet-index" } +reqwest = { version = "0.12", features = ["json"] } serde = { version = "1.0.217", features = ["derive"] } serde_json = "1.0.134" tokio = { version = "1.42.0", features = ["full"] } diff --git a/parakeet/src/hydration/posts.rs b/parakeet/src/hydration/posts.rs index 67c289df..fd430f8e 100644 --- a/parakeet/src/hydration/posts.rs +++ b/parakeet/src/hydration/posts.rs @@ -207,6 +207,7 @@ impl StatefulHydrator<'_> { post, reply, reason: None, + feed_context: None, }, )) }) diff --git a/parakeet/src/main.rs b/parakeet/src/main.rs index 6ee17ef4..db02f1ba 100644 --- a/parakeet/src/main.rs +++ b/parakeet/src/main.rs @@ -19,6 +19,7 @@ mod xrpc; pub struct GlobalState { pub pool: Pool, pub dataloaders: Arc, + pub resolver: Arc, pub index_client: parakeet_index::Client, pub jwt: Arc, pub cdn: Arc, @@ -51,13 +52,13 @@ async fn main() -> eyre::Result<()> { pool.clone(), index_client.clone(), )); - let resolver = did_resolver::Resolver::new(did_resolver::ResolverOpts { + let resolver = Arc::new(did_resolver::Resolver::new(did_resolver::ResolverOpts { plc_directory: conf.plc_directory, ..Default::default() - })?; + })?); let jwt = Arc::new(xrpc::jwt::JwtVerifier::new( conf.service.did.clone(), - resolver, + resolver.clone(), )); let cdn = Arc::new(xrpc::cdn::BskyCdn::new(conf.cdn.base, conf.cdn.video_base)); @@ -82,6 +83,7 @@ async fn main() -> eyre::Result<()> { .with_state(GlobalState { pool, dataloaders, + resolver, index_client, jwt, cdn, diff --git a/parakeet/src/xrpc/app_bsky/feed/posts.rs b/parakeet/src/xrpc/app_bsky/feed/posts.rs index 57469eb9..38af06de 100644 --- a/parakeet/src/xrpc/app_bsky/feed/posts.rs +++ b/parakeet/src/xrpc/app_bsky/feed/posts.rs @@ -5,19 +5,152 @@ use crate::xrpc::extract::{AtpAcceptLabelers, AtpAuth}; use crate::xrpc::{check_actor_status, datetime_cursor, get_actor_did, normalise_at_uri}; use crate::GlobalState; use axum::extract::{Query, State}; +use axum::http::StatusCode; use axum::Json; use axum_extra::extract::Query as ExtraQuery; +use axum_extra::headers::authorization::Bearer; +use axum_extra::headers::Authorization; +use axum_extra::TypedHeader; +use chrono::prelude::*; use diesel::prelude::*; -use diesel_async::RunQueryDsl; +use diesel_async::{AsyncPgConnection, RunQueryDsl}; use lexica::app_bsky::actor::ProfileView; use lexica::app_bsky::feed::{ - FeedViewPost, PostView, ThreadViewPost, ThreadViewPostType, ThreadgateView, + FeedReasonRepost, FeedSkeletonResponse, FeedViewPost, FeedViewPostReason, PostView, + SkeletonReason, ThreadViewPost, ThreadViewPostType, ThreadgateView, }; use parakeet_db::schema; +use reqwest::Url; use serde::{Deserialize, Serialize}; use std::collections::HashMap; -// TODO: getFeed: once we get auth! +const FEEDGEN_SERVICE_ID: &str = "#bsky_fg"; + +#[derive(Debug, Serialize)] +pub struct FeedRes { + #[serde(skip_serializing_if = "Option::is_none")] + cursor: Option, + feed: Vec, +} + +#[derive(Debug, Deserialize)] +pub struct GetFeedQuery { + pub feed: String, + pub limit: Option, + pub cursor: Option, +} + +pub async fn get_feed( + State(state): State, + // we have to use Bearer because the tokens come with `aud` set to the feedgen did. + AtpAcceptLabelers(labelers): AtpAcceptLabelers, + maybe_tok: Option>>, + Query(query): Query, +) -> XrpcResult> { + let mut conn = state.pool.get().await?; + + // first, look up the feedgen + let service_did: String = schema::feedgens::table + .select(schema::feedgens::service_did) + .find(&query.feed) + .get_result(&mut conn) + .await?; + + // resolve the did + let did_doc = match state.resolver.resolve_did(&service_did).await { + Ok(Some(did_doc)) => did_doc, + Ok(None) => return Err(Error::invalid_request(None)), + Err(err) => { + tracing::error!( + feedgen = service_did, + "failed to resolve feedgen service did: {err}" + ); + return Err(Error::invalid_request(None)); + } + }; + + // find the service + let Some(service) = did_doc.find_service_by_id(FEEDGEN_SERVICE_ID) else { + tracing::error!( + feedgen = service_did, + "DID doc didn't contain BskyFeedGenerator service" + ); + return Err(Error::invalid_request(None)); + }; + + let endpoint = service.service_endpoint.clone(); + let skeleton = get_feed_skeleton( + &query.feed, + &endpoint, + maybe_tok.as_ref(), + query.limit, + query.cursor, + ) + .await?; + + let maybe_auth = match maybe_tok { + Some(hdr) => { + match state + .jwt + .resolve_and_verify_jwt(hdr.token(), Some(&service_did)) + .await + { + Some(claims) => match &state.did_allowlist { + Some(allowlist) if !allowlist.contains(&claims.iss) => { + return Err(Error::new( + StatusCode::FORBIDDEN, + "forbidden".to_string(), + None, + )); + } + _ => Some(AtpAuth(claims.iss)), + }, + None => None, + } + } + None => None, + }; + + let hyd = StatefulHydrator::new(&state.dataloaders, &state.cdn, &labelers, maybe_auth); + + let at_uris = skeleton.feed.iter().map(|v| v.post.clone()).collect(); + let repost_skeleton = skeleton + .feed + .iter() + .filter_map(|v| match &v.reason { + Some(SkeletonReason::Repost { repost }) => Some(repost.clone()), + _ => None, + }) + .collect::>(); + + let mut posts = hyd.hydrate_feed_posts(at_uris).await; + let mut repost_data = get_skeleton_repost_data(&mut conn, &hyd, repost_skeleton).await; + + let feed = skeleton + .feed + .into_iter() + .filter_map(|item| { + let mut post = posts.remove(&item.post)?; + let reason = match item.reason { + Some(SkeletonReason::Repost { repost }) => { + repost_data.remove(&repost).map(FeedViewPostReason::Repost) + } + Some(SkeletonReason::Pin {}) => Some(FeedViewPostReason::Pin), + _ => None, + }; + + post.reason = reason; + post.feed_context = item.feed_context; + + Some(post) + }) + .collect(); + + Ok(Json(FeedRes { + cursor: skeleton.cursor, + feed, + })) +} #[derive(Debug, Deserialize)] #[serde(rename_all = "snake_case")] @@ -47,13 +180,6 @@ pub struct GetAuthorFeedQuery { pub include_pins: bool, } -#[derive(Debug, Serialize)] -pub struct FeedRes { - #[serde(skip_serializing_if = "Option::is_none")] - cursor: Option, - feed: Vec, -} - pub async fn get_author_feed( State(state): State, AtpAcceptLabelers(labelers): AtpAcceptLabelers, @@ -460,3 +586,82 @@ fn embed_type_filter<'a>(filter: &'a [&'a str]) -> _ { .eq_any(filter) .or(schema::posts::embed_subtype.eq_any(filter)) } + +async fn get_feed_skeleton( + feed: &str, + service: &str, + maybe_tok: Option<&TypedHeader>>, + limit: Option, + cursor: Option, +) -> XrpcResult { + let mut params = vec![("feed", feed.to_string())]; + + if let Some(cursor) = cursor { + params.push(("cursor", cursor)); + } + if let Some(limit) = limit { + params.push(("limit", limit.to_string())); + } + let url = Url::parse_with_params( + &format!("{service}/xrpc/app.bsky.feed.getFeedSkeleton"), + params, + ) + .unwrap(); + + let mut req = reqwest::Client::new().get(url); + if let Some(auth) = maybe_tok { + req = req.bearer_auth(auth.token()); + } + + match req.send().await { + Ok(skeleton) => match skeleton.json().await { + Ok(skeleton) => Ok(skeleton), + Err(err) => { + tracing::error!("Failed to parse feed skeleton: {err}"); + Err(Error::server_error(Some("Failed to fetch feed skeleton"))) + } + }, + Err(err) => { + tracing::error!("Failed to fetch feed skeleton: {err}"); + Err(Error::server_error(Some("Failed to fetch feed skeleton"))) + } + } +} + +async fn get_skeleton_repost_data<'a>( + conn: &mut AsyncPgConnection, + hyd: &StatefulHydrator<'a>, + reposts: Vec, +) -> HashMap { + let Ok(repost_data) = schema::records::table + .select(( + schema::records::at_uri, + schema::records::did, + schema::records::indexed_at, + )) + .filter(schema::records::at_uri.eq_any(&reposts)) + .get_results::<(String, String, NaiveDateTime)>(conn) + .await + else { + return HashMap::new(); + }; + + let profiles = repost_data.iter().map(|(_, did, _)| did.clone()).collect(); + let profiles = hyd.hydrate_profiles_basic(profiles).await; + + repost_data + .into_iter() + .filter_map(|(uri, did, indexed_at)| { + let by = profiles.get(&did).cloned()?; + + let repost = FeedReasonRepost { + by, + uri: Some(uri.clone()), + cid: None, // okay, we do have this, but the app doesn't seem to be bothered about not setting it. + indexed_at: indexed_at.and_utc(), + }; + + Some((uri, repost)) + }) + .collect() +} diff --git a/parakeet/src/xrpc/app_bsky/mod.rs b/parakeet/src/xrpc/app_bsky/mod.rs index 63db453e..101810e6 100644 --- a/parakeet/src/xrpc/app_bsky/mod.rs +++ b/parakeet/src/xrpc/app_bsky/mod.rs @@ -14,6 +14,7 @@ pub fn routes() -> Router { .route("/app.bsky.feed.getActorFeeds", get(feed::feedgen::get_actor_feeds)) .route("/app.bsky.feed.getActorLikes", get(feed::likes::get_actor_likes)) .route("/app.bsky.feed.getAuthorFeed", get(feed::posts::get_author_feed)) + .route("/app.bsky.feed.getFeed", get(feed::posts::get_feed)) .route("/app.bsky.feed.getLikes", get(feed::likes::get_likes)) .route("/app.bsky.feed.getListFeed", get(feed::posts::get_list_feed)) .route("/app.bsky.feed.getPostThread", get(feed::posts::get_post_thread)) diff --git a/parakeet/src/xrpc/extract.rs b/parakeet/src/xrpc/extract.rs index 242e09d0..bffc984b 100644 --- a/parakeet/src/xrpc/extract.rs +++ b/parakeet/src/xrpc/extract.rs @@ -81,7 +81,7 @@ impl FromRequestParts for AtpAuth { .map_err(|err| (StatusCode::BAD_REQUEST, err.to_string()))? .ok_or((StatusCode::UNAUTHORIZED, "missing JWT".to_string()))?; - match state.jwt.resolve_and_verify_jwt(hdr.token()).await { + match state.jwt.resolve_and_verify_jwt(hdr.token(), None).await { Some(claims) => match &state.did_allowlist { Some(allowlist) if !allowlist.contains(&claims.iss) => { Err((StatusCode::FORBIDDEN, "forbidden".to_string())) @@ -110,7 +110,7 @@ impl OptionalFromRequestParts for AtpAuth { return Ok(None); }; - match state.jwt.resolve_and_verify_jwt(hdr.token()).await { + match state.jwt.resolve_and_verify_jwt(hdr.token(), None).await { Some(claims) => match &state.did_allowlist { Some(allowlist) if !allowlist.contains(&claims.iss) => { Err((StatusCode::FORBIDDEN, "forbidden".to_string())) diff --git a/parakeet/src/xrpc/jwt.rs b/parakeet/src/xrpc/jwt.rs index 78d973d0..772a5da1 100644 --- a/parakeet/src/xrpc/jwt.rs +++ b/parakeet/src/xrpc/jwt.rs @@ -2,7 +2,7 @@ use did_resolver::Resolver; use jsonwebtoken::{Algorithm, DecodingKey, Validation}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; -use std::sync::LazyLock; +use std::sync::{Arc, LazyLock}; use tokio::sync::RwLock; static DUMMY_KEY: LazyLock = LazyLock::new(|| DecodingKey::from_secret(&[])); @@ -25,12 +25,12 @@ pub struct Claims { pub struct JwtVerifier { aud: String, - resolver: Resolver, + resolver: Arc, key_cache: RwLock>, } impl JwtVerifier { - pub fn new(aud: String, resolver: Resolver) -> Self { + pub fn new(aud: String, resolver: Arc) -> Self { JwtVerifier { aud, resolver, @@ -38,7 +38,7 @@ impl JwtVerifier { } } - pub async fn resolve_and_verify_jwt(&self, token: &str) -> Option { + pub async fn resolve_and_verify_jwt(&self, token: &str, aud: Option<&str>) -> Option { // first we need to decode without verifying, to get iss. let unsafe_data = jsonwebtoken::decode::(token, &DUMMY_KEY, &NO_VERIFY).ok()?; let unsafe_iss = unsafe_data.claims.iss; @@ -52,7 +52,8 @@ impl JwtVerifier { None => self.resolve_key(&unsafe_iss).await?, }; - self.verify_jwt_multibase_with_alg(token, &multibase_key, unsafe_data.header.alg) + let aud = aud.unwrap_or(&self.aud); + self.verify_jwt_multibase_with_alg(token, &multibase_key, unsafe_data.header.alg, aud) } async fn resolve_key(&self, did: &str) -> Option { @@ -73,7 +74,7 @@ impl JwtVerifier { pub fn verify_jwt_multibase(&self, token: &str, multibase_key: &str) -> Option { let alg = jsonwebtoken::decode_header(token).ok()?.alg; - self.verify_jwt_multibase_with_alg(token, multibase_key, alg) + self.verify_jwt_multibase_with_alg(token, multibase_key, alg, &self.aud) } pub fn verify_jwt_multibase_with_alg( @@ -81,6 +82,7 @@ impl JwtVerifier { token: &str, multibase_key: &str, alg: Algorithm, + aud: &str, ) -> Option { // decode the multibase key let (_, key) = multibase::decode(multibase_key).ok()?; @@ -88,7 +90,7 @@ impl JwtVerifier { let key = DecodingKey::from_ec_der(&key[2..]); let mut validation = Validation::new(alg); - validation.set_audience(&[&self.aud]); + validation.set_audience(&[&aud]); let decoded = jsonwebtoken::decode::(token, &key, &validation).ok()?;