You are a constitutional council ranking individual git commits for ownership allocation. Compare these two commits. Decide which contributed more lasting value to the project. Judge substance, not spectacle: - Prefer correct, lasting design and real bugfixes over churn, formatting, renames, or generated noise. - Prefer clarity and necessity over sheer line count. A small precise change can beat a large diffuse one. - Do not favor a side merely because its patch is longer or noisier. - Weight what the change does for the project, not the contributor's name. Return ONLY a JSON object: {"winner": "A" or "B", "ratio": "N:M", "explanation": "..."} The explanation must cite concrete differences in the patches (1-3 sentences). Side A — contributor: tommy-mor Side A — commit message: [52f5c51c] Add Reddit OAuth linking and make UUID the only account identity. OAuth providers only attach to a session UUID (first link creates the principal); linked providers stay private on the account page. Co-authored-by: Cursor Side A — unified diff (full patch): diff --git a/AGENTS.md b/AGENTS.md index 6e0fd8ebb65d665c9c1438e3275971d62b98fd95..e9cc3173dbeb21ad0fc090ca7b407b027c7820a9 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -35,8 +35,11 @@ Environment variables (defaults in `server/src/state.rs`): - `SORTER2_DATA_DIR` — default `./data` (created on startup) - `SORTER2_EVENT_LOG` — default `{data_dir}/events.jsonl` - `SORTER2_BASE_URL` — public origin (also drives Secure cookies when `https://`) -- `GITHUB_CLIENT_ID` / `GITHUB_CLIENT_SECRET` — GitHub OAuth (optional; login disabled if unset) -- `SORTER2_ALLOW_MOCK_OAUTH=1` — allow `mock_user` on `/auth/github` (tests only) +- `GITHUB_CLIENT_ID` / `GITHUB_CLIENT_SECRET` — GitHub OAuth linking (optional) +- `REDDIT_CLIENT_ID` / `REDDIT_CLIENT_SECRET` (or `REDDIT_APP_*`) — Reddit API import + OAuth linking (optional) +- `SORTER2_ALLOW_MOCK_OAUTH=1` — allow `mock_user` on `/auth/github` and `/auth/reddit` (tests only) + +Identity: UUID is canonical. OAuth providers only *link* to a UUID (first link creates the principal). Linked providers are private to the account owner. Health check: `GET /healthz` → `ok`. diff --git a/server/src/auth/mod.rs b/server/src/auth/mod.rs index 5906f93b13853421e96a3c37bc9d8202a47842bf..c706ae8045a811e5941f5f6c72da88f42a403a82 100644 --- a/server/src/auth/mod.rs +++ b/server/src/auth/mod.rs @@ -1,4 +1,8 @@ -//! GitHub OAuth login, session cookies, and vote actor resolution. +//! OAuth linking, session cookies, and vote actor resolution. +//! +//! Canonical identity is a UUID. OAuth providers only *link* to that UUID +//! (first link creates the principal; later links attach while logged in). +//! Which providers are linked is private to the account owner. pub mod config; pub mod identity; @@ -22,7 +26,9 @@ use crate::{ form_template::template_json_compact, html::layout, state::AppState, - storage_schema::{oauth_link_owner, pseudonym_owner, Store, StoreFields}, + storage_schema::{ + linked_providers_for_uuid, oauth_link_owner, pseudonym_owner, Store, StoreFields, + }, ui_action::UI_RPC_FIELD, }; @@ -53,10 +59,12 @@ fn new_actor_uuid() -> String { pub struct LoginQuery { #[serde(default)] pub return_to: Option, + #[serde(default)] + pub error: Option, } #[derive(Debug, Deserialize)] -pub struct GitHubStartQuery { +pub struct OAuthStartQuery { #[serde(default)] pub return_to: Option, #[serde(default)] @@ -72,15 +80,22 @@ fn return_from_query_or_jar(jar: &CookieJar, query: Option<&str>) -> String { .unwrap_or_else(|| "/".to_string()) } -fn oauth_providers(base_url: &str, return_to: &str) -> Vec<(&'static str, String)> { +/// Available OAuth link targets: `(provider_key, label, start_href)`. +fn oauth_providers(base_url: &str, return_to: &str) -> Vec<(&'static str, &'static str, String)> { let mut out = Vec::new(); + let enc = urlencoding::encode(return_to); if oauth::GitHubConfig::from_env(base_url).is_some() { out.push(( - "GitHub", - format!( - "/auth/github?return_to={}", - urlencoding::encode(return_to) - ), + "github", + oauth::provider_label("github"), + format!("/auth/github?return_to={enc}"), + )); + } + if oauth::RedditConfig::from_env(base_url).is_some() { + out.push(( + "reddit", + oauth::provider_label("reddit"), + format!("/auth/reddit?return_to={enc}"), )); } out @@ -125,23 +140,41 @@ fn alias_claim_forms(return_to: &str, submit_label: &str) -> Result Markup { +fn login_error_message(code: Option<&str>) -> Option<&'static str> { + match code { + Some("oauth_taken") => { + Some("that OAuth account is already linked to a different sorter2 account") + } + Some("oauth_failed") => Some("OAuth failed — try again"), + _ => None, + } +} + +fn signed_out_body( + providers: &[(&str, &str, String)], + error: Option<&str>, +) -> Markup { html! { main class="panel login-page" { section class="login-section" { h1 { "sign in" } - p class="muted" { "link an account to vote under a lasting alias" } + p class="muted" { + "link an OAuth account to create your identity, then claim an alias to vote" + } + @if let Some(msg) = login_error_message(error) { + p class="alias-bad" data-testid="login-error" { (msg) } + } @if providers.is_empty() { p class="muted" { - "OAuth is not configured. Set GITHUB_CLIENT_ID and GITHUB_CLIENT_SECRET." + "OAuth is not configured. Set GitHub and/or Reddit client credentials." } } @else { ul class="oauth-provider-list" { - @for (name, href) in providers { + @for (key, label, href) in providers { li { a href=(href) class="btn-primary oauth-provider" - data-testid=(format!("oauth-{}", name.to_lowercase())) { - (format!("Continue with {name}")) + data-testid=(format!("oauth-{key}")) { + (format!("Link {label}")) } } } @@ -156,7 +189,10 @@ fn signed_out_body(providers: &[(&str, String)]) -> Markup { fn account_body( actor: &session::SessionActor, aliases: &[String], - providers: &[(&str, String)], + // Provider keys already linked to this UUID (private). + linked: &[String], + // Providers available to link: not yet attached. + unlinkable: &[(&str, &str, String)], claim_forms: Markup, ) -> Markup { let current = actor.pseudonym.trim(); @@ -212,16 +248,29 @@ fn account_body( (claim_forms) } - @if !providers.is_empty() { - section class="login-section" { - h2 { "linked sign-in" } - p class="muted small" { "sign in again with the same provider to return to this account" } + section class="login-section" { + h2 { "linked sign-in" } + p class="muted small" { + "private to you — linking more providers raises trust weight without publishing which accounts you use" + } + @if linked.is_empty() { + p class="muted" data-testid="linked-providers-empty" { "none yet" } + } @else { + ul class="linked-provider-list" data-testid="linked-providers" { + @for key in linked { + li data-testid=(format!("linked-{key}")) { + (oauth::provider_label(key)) + } + } + } + } + @if !unlinkable.is_empty() { ul class="oauth-provider-list" { - @for (name, href) in providers { + @for (key, label, href) in unlinkable { li { a href=(href) class="btn-secondary oauth-provider" - data-testid=(format!("oauth-relink-{}", name.to_lowercase())) { - (format!("Re-link {name}")) + data-testid=(format!("oauth-link-{key}")) { + (format!("Link {label}")) } } } @@ -243,12 +292,21 @@ fn account_body( fn login_body( session: Option<&session::SessionActor>, aliases: &[String], - providers: &[(&str, String)], + linked: &[String], + providers: &[(&str, &str, String)], claim_forms: Option, + error: Option<&str>, ) -> Markup { match (session, claim_forms) { - (Some(actor), Some(forms)) => account_body(actor, aliases, providers, forms), - _ => signed_out_body(providers), + (Some(actor), Some(forms)) => { + let unlinkable: Vec<_> = providers + .iter() + .filter(|(key, _, _)| !linked.iter().any(|p| p == key)) + .cloned() + .collect(); + account_body(actor, aliases, linked, &unlinkable, forms) + } + _ => signed_out_body(providers, error), } } @@ -268,6 +326,10 @@ pub async fn login_page( .as_ref() .map(|s| alias_list(db, &s.uuid)) .unwrap_or_default(); + let linked = session + .as_ref() + .map(|s| linked_providers_for_uuid(db, &s.uuid).unwrap_or_default()) + .unwrap_or_default(); let providers = oauth_providers(&base_url_from_env(state.cfg.port), &return_to); let claim_forms = if session.is_some() { @@ -282,7 +344,14 @@ pub async fn login_page( } else { "login · sorter2" }, - login_body(session.as_ref(), &aliases, &providers, claim_forms), + login_body( + session.as_ref(), + &aliases, + &linked, + &providers, + claim_forms, + query.error.as_deref(), + ), state.views.get_views("/login"), session .as_ref() @@ -302,7 +371,6 @@ pub async fn alias_page( let db = state.projection_store.db(); let session = session::load_valid_session(db, &session_id).ok_or(StatusCode::UNAUTHORIZED)?; if session::session_has_pseudonym(&session) { - // Already onboarded — manage aliases on the account page. return Ok(Redirect::to("/login").into_response()); } @@ -331,7 +399,7 @@ pub async fn alias_page( pub async fn github_start( State(state): State, jar: CookieJar, - Query(query): Query, + Query(query): Query, ) -> Result { let cfg = oauth::GitHubConfig::from_env(&base_url_from_env(state.cfg.port)) .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; @@ -342,7 +410,28 @@ pub async fn github_start( } else { None }; - let url = oauth::authorize_url(&cfg, &state_token, mock_user); + let url = oauth::github_authorize_url(&cfg, &state_token, mock_user); + let jar = jar + .add(session::oauth_state_cookie_value(&state_token)) + .add(session::auth_return_cookie_value(&return_to)); + Ok((jar, Redirect::temporary(&url)).into_response()) +} + +pub async fn reddit_start( + State(state): State, + jar: CookieJar, + Query(query): Query, +) -> Result { + let cfg = oauth::RedditConfig::from_env(&base_url_from_env(state.cfg.port)) + .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; + let return_to = return_from_query_or_jar(&jar, query.return_to.as_deref()); + let state_token = session::new_oauth_state(); + let mock_user = if config::mock_oauth_allowed() { + query.mock_user.as_deref() + } else { + None + }; + let url = oauth::reddit_authorize_url(&cfg, &state_token, mock_user); let jar = jar .add(session::oauth_state_cookie_value(&state_token)) .add(session::auth_return_cookie_value(&return_to)); @@ -355,6 +444,13 @@ pub struct OAuthCallbackQuery { pub state: String, } +/// Link `provider:provider_id` to a UUID. +/// +/// - Logged in + new provider → attach to session UUID +/// - Logged in + already ours → no-op +/// - Logged in + owned by someone else → conflict +/// - Logged out + known link → resume that UUID +/// - Logged out + unknown → create principal + first link async fn finish_oauth_login( state: &AppState, jar: CookieJar, @@ -364,11 +460,41 @@ async fn finish_oauth_login( let db = state.projection_store.db(); let return_to = return_from_query_or_jar(&jar, None); - let uuid = match oauth_link_owner(db, provider, &provider_id) - .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)? - { - Some(existing) => existing, - None => { + let existing_owner = oauth_link_owner(db, provider, &provider_id) + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; + + let session_uuid = session::session_id_from_jar(&jar) + .as_deref() + .and_then(|id| session::load_valid_session(db, id)) + .map(|s| s.uuid); + let linking_while_logged_in = session_uuid.is_some(); + + let uuid = match (session_uuid, existing_owner) { + (Some(session_uuid), Some(owner)) if owner == session_uuid => session_uuid, + (Some(_), Some(_)) => { + return Ok(( + jar.add(session::clear_oauth_state_cookie()), + "/login?error=oauth_taken".into(), + )); + } + (Some(session_uuid), None) => { + let ts = now_ms(); + state + .append_identity_events(vec![Event::OauthLinked { + uuid: session_uuid.clone(), + provider: provider.to_string(), + provider_id, + ts, + }]) + .await + .map_err(|e| { + tracing::warn!(err = %e, "oauth link append failed"); + StatusCode::INTERNAL_SERVER_ERROR + })?; + session_uuid + } + (None, Some(owner)) => owner, + (None, None) => { let uuid = new_actor_uuid(); let ts = now_ms(); state @@ -409,6 +535,9 @@ async fn finish_oauth_login( "/login/alias?return_to={}", urlencoding::encode(&return_to) ) + } else if linking_while_logged_in { + // Additional link while already in an account → stay on account page. + "/login".to_string() } else { return_to }; @@ -434,22 +563,57 @@ pub async fn github_callback( .build() .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; - let token = oauth::exchange_code(&client, &cfg, &query.code) + let token = oauth::github_exchange_code(&client, &cfg, &query.code) .await .map_err(|e| { tracing::warn!(err = %e, "github oauth token exchange failed"); StatusCode::BAD_GATEWAY })?; - let user = oauth::fetch_user(&client, &cfg.api_base, &token) + let user = oauth::github_fetch_user(&client, &cfg.api_base, &token) .await .map_err(|e| { tracing::warn!(err = %e, "github user fetch failed"); StatusCode::BAD_GATEWAY })?; - let provider = "github"; - let provider_id = oauth::provider_id(&user); - let (jar, dest) = finish_oauth_login(&state, jar, provider, provider_id).await?; + let (jar, dest) = + finish_oauth_login(&state, jar, "github", oauth::github_provider_id(&user)).await?; + Ok((jar, Redirect::to(&dest)).into_response()) +} + +pub async fn reddit_callback( + State(state): State, + jar: CookieJar, + Query(query): Query, +) -> Result { + let cfg = oauth::RedditConfig::from_env(&base_url_from_env(state.cfg.port)) + .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; + + let expected_state = session::oauth_state_from_jar(&jar).ok_or(StatusCode::BAD_REQUEST)?; + if expected_state != query.state { + return Err(StatusCode::BAD_REQUEST); + } + + let client = Client::builder() + .timeout(std::time::Duration::from_secs(15)) + .build() + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; + + let token = oauth::reddit_exchange_code(&client, &cfg, &query.code) + .await + .map_err(|e| { + tracing::warn!(err = %e, "reddit oauth token exchange failed"); + StatusCode::BAD_GATEWAY + })?; + let user = oauth::reddit_fetch_user(&client, &cfg, &token) + .await + .map_err(|e| { + tracing::warn!(err = %e, "reddit user fetch failed"); + StatusCode::BAD_GATEWAY + })?; + + let (jar, dest) = + finish_oauth_login(&state, jar, "reddit", oauth::reddit_provider_id(&user)).await?; Ok((jar, Redirect::to(&dest)).into_response()) } diff --git a/server/src/auth/oauth.rs b/server/src/auth/oauth.rs index b80078fee0c5acd905b9125d29454e46c4e0d066..369b1527d9032cf312a07819279e0f988ef4a36e 100644 --- a/server/src/auth/oauth.rs +++ b/server/src/auth/oauth.rs @@ -1,8 +1,15 @@ -//! GitHub OAuth (raw reqwest, same style as reddit.rs). +//! OAuth providers (GitHub + Reddit). Provider accounts only *link* to a UUID; +//! the UUID is the canonical identity. Which providers are linked is private. use reqwest::Client; use serde::Deserialize; +use crate::reddit::{ + default_user_agent, reddit_oauth_api_base, reddit_oauth_token_base, +}; + +// ── GitHub ────────────────────────────────────────────────────────────────── + #[derive(Debug, Clone)] pub struct GitHubConfig { pub client_id: String, @@ -40,7 +47,7 @@ impl GitHubConfig { } #[derive(Debug, Deserialize)] -struct TokenResponse { +struct GitHubTokenResponse { access_token: String, } @@ -50,7 +57,7 @@ pub struct GitHubUser { pub login: String, } -pub fn authorize_url(cfg: &GitHubConfig, state: &str, mock_user: Option<&str>) -> String { +pub fn github_authorize_url(cfg: &GitHubConfig, state: &str, mock_user: Option<&str>) -> String { let mut url = format!( "{}/login/oauth/authorize?client_id={}&redirect_uri={}&scope=read:user&state={}", cfg.oauth_base.trim_end_matches('/'), @@ -65,7 +72,7 @@ pub fn authorize_url(cfg: &GitHubConfig, state: &str, mock_user: Option<&str>) - url } -pub async fn exchange_code( +pub async fn github_exchange_code( client: &Client, cfg: &GitHubConfig, code: &str, @@ -90,14 +97,14 @@ pub async fn exchange_code( return Err(format!("github token HTTP {}", resp.status())); } - let body: TokenResponse = resp + let body: GitHubTokenResponse = resp .json() .await .map_err(|e| format!("github token parse failed: {e}"))?; Ok(body.access_token) } -pub async fn fetch_user( +pub async fn github_fetch_user( client: &Client, api_base: &str, access_token: &str, @@ -120,10 +127,149 @@ pub async fn fetch_user( .map_err(|e| format!("github user parse failed: {e}")) } -pub fn provider_id(user: &GitHubUser) -> String { +pub fn github_provider_id(user: &GitHubUser) -> String { user.id.to_string() } +// ── Reddit ────────────────────────────────────────────────────────────────── + +#[derive(Debug, Clone)] +pub struct RedditConfig { + pub client_id: String, + pub client_secret: String, + pub redirect_uri: String, + /// Host for `/api/v1/authorize` (www.reddit.com in production). + pub authorize_base: String, + /// Host for `POST /api/v1/access_token`. + pub token_base: String, + /// Host for bearer `GET /api/v1/me` (oauth.reddit.com). + pub api_base: String, + pub user_agent: String, +} + +/// Authorize page base; defaults to the same host as token POSTs. +pub fn reddit_authorize_base() -> String { + std::env::var("REDDIT_OAUTH_AUTHORIZE_BASE") + .or_else(|_| std::env::var("REDDIT_OAUTH_BASE")) + .unwrap_or_else(|_| "https://www.reddit.com".into()) +} + +impl RedditConfig { + pub fn from_env(base_url: &str) -> Option { + let client_id = std::env::var("REDDIT_CLIENT_ID") + .or_else(|_| std::env::var("REDDIT_APP_ID")) + .ok()?; + let client_secret = std::env::var("REDDIT_CLIENT_SECRET") + .or_else(|_| std::env::var("REDDIT_APP_SECRET")) + .ok()?; + if client_id.is_empty() || client_secret.is_empty() { + return None; + } + let base = base_url.trim_end_matches('/'); + Some(Self { + client_id, + client_secret, + redirect_uri: format!("{base}/auth/reddit/callback"), + authorize_base: reddit_authorize_base(), + token_base: reddit_oauth_token_base(), + api_base: reddit_oauth_api_base(), + user_agent: default_user_agent(), + }) + } +} + +#[derive(Debug, Deserialize)] +struct RedditTokenResponse { + access_token: String, +} + +#[derive(Debug, Deserialize)] +pub struct RedditUser { + /// Stable id (`t2_…`); never use `name` as identity. + pub id: String, + pub name: String, +} + +pub fn reddit_authorize_url(cfg: &RedditConfig, state: &str, mock_user: Option<&str>) -> String { + let mut url = format!( + "{}/api/v1/authorize?client_id={}&response_type=code&state={}&redirect_uri={}&duration=temporary&scope=identity", + cfg.authorize_base.trim_end_matches('/'), + urlencoding::encode(&cfg.client_id), + urlencoding::encode(state), + urlencoding::encode(&cfg.redirect_uri), + ); + if let Some(user) = mock_user { + url.push_str("&mock_user="); + url.push_str(&urlencoding::encode(user)); + } + url +} + +pub async fn reddit_exchange_code( + client: &Client, + cfg: &RedditConfig, + code: &str, +) -> Result { + let resp = client + .post(format!( + "{}/api/v1/access_token", + cfg.token_base.trim_end_matches('/') + )) + .header("User-Agent", &cfg.user_agent) + .basic_auth(&cfg.client_id, Some(&cfg.client_secret)) + .form(&[ + ("grant_type", "authorization_code"), + ("code", code), + ("redirect_uri", cfg.redirect_uri.as_str()), + ]) + .send() + .await + .map_err(|e| format!("reddit token request failed: {e}"))?; + + if !resp.status().is_success() { + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + return Err(format!("reddit token HTTP {status}: {body}")); + } + + let body: RedditTokenResponse = resp + .json() + .await + .map_err(|e| format!("reddit token parse failed: {e}"))?; + Ok(body.access_token) +} + +pub async fn reddit_fetch_user( + client: &Client, + cfg: &RedditConfig, + access_token: &str, +) -> Result { + let resp = client + .get(format!( + "{}/api/v1/me", + cfg.api_base.trim_end_matches('/') + )) + .header("User-Agent", &cfg.user_agent) + .bearer_auth(access_token) + .send() + .await + .map_err(|e| format!("reddit user request failed: {e}"))?; + + if !resp.status().is_success() { + return Err(format!("reddit user HTTP {}", resp.status())); + } + + resp.json() + .await + .map_err(|e| format!("reddit user parse failed: {e}")) +} + +pub fn reddit_provider_id(user: &RedditUser) -> String { + user.id.clone() +} + +// ── Shared helpers ────────────────────────────────────────────────────────── + pub fn validate_pseudonym(raw: &str) -> Result { let trimmed = raw.trim(); if trimmed.is_empty() { @@ -144,3 +290,12 @@ pub fn validate_pseudonym(raw: &str) -> Result { pub fn sanitize_pseudonym(login: &str) -> String { validate_pseudonym(login).unwrap_or_else(|_| "user".to_string()) } + +/// Display name for a provider key (`github` → `GitHub`). Never show provider ids. +pub fn provider_label(provider: &str) -> &'static str { + match provider { + "github" => "GitHub", + "reddit" => "Reddit", + _ => "OAuth", + } +} diff --git a/server/src/lib.rs b/server/src/lib.rs index 84f2565b105fb302241b949af64bd5e49916eab2..0d948af1457504fdbd34d0b14261f65ac58da0ec 100644 --- a/server/src/lib.rs +++ b/server/src/lib.rs @@ -47,6 +47,8 @@ pub fn create_app(state: AppState) -> Router { .route("/login/alias", get(crate::auth::alias_page)) .route("/auth/github", get(crate::auth::github_start)) .route("/auth/github/callback", get(crate::auth::github_callback)) + .route("/auth/reddit", get(crate::auth::reddit_start)) + .route("/auth/reddit/callback", get(crate::auth::reddit_callback)) .route("/auth/logout", post(crate::auth::logout)) .route("/auth/switch", post(crate::auth::switch_pseudonym)) .route("/ui", post(crate::api::ui_html::post_ui_html)) diff --git a/server/src/projection_apply.rs b/server/src/projection_apply.rs index 4fb4b48ee0a53b7f673fae0c9731eed5acc33332..f50458c2ff3c447a3cfd28adb298e6883a2da4e9 100644 --- a/server/src/projection_apply.rs +++ b/server/src/projection_apply.rs @@ -4,6 +4,8 @@ //! child links, recent-vote appends) plus a cursor advance, all committed in //! one atomic `DisableWal` batch. +use std::collections::HashMap; + use crate::{ event_log::EventLogError, events::{Event, EventRecord}, @@ -44,6 +46,8 @@ pub fn apply_records( let db = projection_store.db(); let mut batch = db.batch(); let mut last_seq = 0u64; + // Weight reads must see earlier writes in this same batch. + let mut pending_weights: HashMap = HashMap::new(); for record in records { match &record.event { @@ -85,6 +89,7 @@ pub fn apply_records( ensure_path_writes(&mut batch, &parsed); } Event::PrincipalCreated { uuid, .. } => { + pending_weights.insert(uuid.clone(), BASE_TRUST_WEIGHT); batch.write( Store::root() .user_weights() @@ -112,17 +117,25 @@ pub fn apply_records( } } else { batch.write(Store::root().oauth_links().key(&link_key).set(uuid)); - let current = Store::root() - .user_weights() - .key(&uuid.clone()) - .get(db) - .map_err(|e| EventLogError::Apply(e.to_string()))? + let current = pending_weights + .get(uuid) + .copied() + .or_else(|| { + Store::root() + .user_weights() + .key(&uuid.clone()) + .get(db) + .ok() + .flatten() + }) .unwrap_or(BASE_TRUST_WEIGHT); + let next = trust_weight_after_link(current); + pending_weights.insert(uuid.clone(), next); batch.write( Store::root() .user_weights() .key(&uuid.clone()) - .set(&trust_weight_after_link(current)), + .set(&next), ); } } @@ -205,6 +218,15 @@ mod tests { ), record( 3, + Event::OauthLinked { + uuid: uuid.into(), + provider: "reddit".into(), + provider_id: "t2_abc".into(), + ts, + }, + ), + record( + 4, Event::PseudonymClaimed { uuid: uuid.into(), pseudonym: "octocat".into(), @@ -219,8 +241,12 @@ mod tests { oauth_link_owner(store.db(), "github", "42").unwrap(), Some(uuid.to_string()) ); + assert_eq!( + crate::storage_schema::linked_providers_for_uuid(store.db(), uuid).unwrap(), + vec!["github".to_string(), "reddit".to_string()] + ); assert_eq!(resolve_actor_uuid(store.db(), "octocat").unwrap(), uuid); - assert_eq!(user_trust_weight(store.db(), uuid).unwrap(), 1.5); + assert_eq!(user_trust_weight(store.db(), uuid).unwrap(), 2.0); let aliases = Store::root() .user_pseudonyms() .key(&uuid.to_string()) diff --git a/server/src/storage_schema.rs b/server/src/storage_schema.rs index b942810afd96d3bf01b7765718303a39be4ef8d2..dac6f8e6080e3f5683c24dbb8861f5a981d456c4 100644 --- a/server/src/storage_schema.rs +++ b/server/src/storage_schema.rs @@ -103,6 +103,25 @@ pub fn oauth_link_owner(db: &Db, provider: &str, provider_id: &str) -> durable:: .get(db) } +/// Provider names linked to a UUID (`github`, `reddit`, …). Private — for the +/// account owner's page only; never expose which providers are linked publicly. +pub fn linked_providers_for_uuid(db: &Db, uuid: &str) -> durable::Result> { + let mut providers = Vec::new(); + for (key, owner) in Store::root().oauth_links().iter(db)? { + if owner != uuid { + continue; + } + let Some((provider, _)) = key.split_once(':') else { + continue; + }; + if !providers.iter().any(|p| p == provider) { + providers.push(provider.to_string()); + } + } + providers.sort(); + Ok(providers) +} + pub const RECENT_VOTES_CAP: u64 = 200; fn id_key(id: &ItemId) -> String { diff --git a/test/support/harness.clj b/test/support/harness.clj index 4505ece5aa193e826ca61c52f1467d252080d66b..741024aab1d14269e225791189633287980ae2e4 100644 --- a/test/support/harness.clj +++ b/test/support/harness.clj @@ -48,12 +48,15 @@ "GITHUB_CLIENT_SECRET" "test-secret" "GITHUB_OAUTH_BASE" (str "http://127.0.0.1:" oauth-port) "GITHUB_API_BASE" (str "http://127.0.0.1:" oauth-port) + ;; Reddit import fixtures + Reddit OAuth on reddit-port. "REDDIT_API_BASE" (str "http://127.0.0.1:" reddit-port) + "REDDIT_CLIENT_ID" "test-reddit" + "REDDIT_CLIENT_SECRET" "test-reddit-secret" "REDDIT_OAUTH_BASE" (str "http://127.0.0.1:" reddit-port) - "REDDIT_CLIENT_ID" "" - "REDDIT_CLIENT_SECRET" "" - "REDDIT_APP_ID" "" - "REDDIT_APP_SECRET" ""})) + "REDDIT_OAUTH_AUTHORIZE_BASE" (str "http://127.0.0.1:" reddit-port) + "REDDIT_OAUTH_TOKEN_BASE" (str "http://127.0.0.1:" reddit-port) + "REDDIT_OAUTH_API_BASE" (str "http://127.0.0.1:" reddit-port) + "REDDIT_USER_AGENT" "web:sorter2-test:v0 (by /u/test)"})) (defn with-auth-servers "Start mock Reddit + mock OAuth + release sorter2-server. diff --git a/test/support/mock_oauth.clj b/test/support/mock_oauth.clj index 909d7a7be8b159ea54a122d69c46af3db3318118..5ba7e9be3cf6d2226648e9609a09ed45f306930f 100644 --- a/test/support/mock_oauth.clj +++ b/test/support/mock_oauth.clj @@ -1,5 +1,5 @@ (ns test.support.mock-oauth - "In-process HTTP stub for GitHub OAuth (authorize, token, /user)." + "In-process HTTP stub for GitHub + Reddit OAuth (authorize, token, user)." (:require [clojure.string :as str]) (:import [com.sun.net.httpserver HttpServer HttpHandler HttpExchange] [java.net InetSocketAddress URLDecoder])) @@ -12,11 +12,14 @@ (URLDecoder/decode (or v "") "UTF-8")))) (str/split query #"&")))) -(defn- parse-mock-user [raw] +(defn- parse-mock-user + "GitHub-style `id:login` (numeric id). Reddit-style `t2_xxx:name`." + [raw] (let [s (or raw "1002:newbie") [id login] (str/split s #":" 2)] - {:id (Long/parseLong id) - :login (or login "newbie")})) + {:id id + :login (or login "newbie") + :numeric? (re-matches #"\d+" id)})) (defn- send-json [^HttpExchange ex status body] (let [bytes (.getBytes body "UTF-8")] @@ -31,9 +34,10 @@ (.sendResponseHeaders ex 302 -1) (.close (.getResponseBody ex))) -(defn- read-form-code [^HttpExchange ex] +(defn- read-form [^HttpExchange ex] (let [body (slurp (.getInputStream ex))] - (query-param body "code"))) + {:code (query-param body "code") + :grant (query-param body "grant_type")})) (defn- bearer-token [^HttpExchange ex] (some-> (.getRequestHeaders ex) @@ -44,8 +48,18 @@ (when (str/starts-with? token "mock:") (parse-mock-user (subs token 5)))) +(defn- authorize-redirect [exchange query] + (let [redirect-uri (query-param query "redirect_uri") + state (query-param query "state") + mock-user (query-param query "mock_user") + user (parse-mock-user mock-user) + code (str "mock:" (:id user) ":" (:login user)) + loc (str redirect-uri "?code=" (java.net.URLEncoder/encode code "UTF-8") + "&state=" (java.net.URLEncoder/encode state "UTF-8"))] + (send-redirect exchange loc))) + (defn start-mock-oauth - "Start mock GitHub OAuth on `port`. Returns a zero-arg `stop` function." + "Start mock GitHub + Reddit OAuth on `port`. Returns a zero-arg `stop` function." [port] (let [server (HttpServer/create (InetSocketAddress. "127.0.0.1" port) 0) handler @@ -53,28 +67,45 @@ (handle [^HttpExchange exchange] (let [uri (.getRequestURI exchange) path (.getPath uri) - query (.getQuery uri)] + query (.getQuery uri) + method (.getRequestMethod exchange)] (cond + ;; GitHub authorize (str/ends-with? path "/login/oauth/authorize") - (let [redirect-uri (query-param query "redirect_uri") - state (query-param query "state") - mock-user (query-param query "mock_user") - user (parse-mock-user mock-user) - code (str "mock:" (:id user) ":" (:login user)) - loc (str redirect-uri "?code=" (java.net.URLEncoder/encode code "UTF-8") - "&state=" (java.net.URLEncoder/encode state "UTF-8"))] - (send-redirect exchange loc)) + (authorize-redirect exchange query) + + ;; Reddit authorize + (str/ends-with? path "/api/v1/authorize") + (authorize-redirect exchange query) - (str/ends-with? path "/login/oauth/access_token") - (let [code (or (read-form-code exchange) "mock:1002:newbie")] + ;; GitHub token + (and (= method "POST") (str/ends-with? path "/login/oauth/access_token")) + (let [code (or (:code (read-form exchange)) "mock:1002:newbie")] (send-json exchange 200 (str "{\"access_token\":\"" code "\",\"token_type\":\"bearer\"}"))) + ;; Reddit token (client_credentials for import + authorization_code for login) + (and (= method "POST") (str/ends-with? path "/api/v1/access_token")) + (let [form (read-form exchange) + grant (or (:grant form) "") + code (or (:code form) "mock:t2_test:redditor")] + (if (= grant "client_credentials") + (send-json exchange 200 "{\"access_token\":\"app-token\",\"token_type\":\"bearer\",\"expires_in\":3600}") + (send-json exchange 200 (str "{\"access_token\":\"" code "\",\"token_type\":\"bearer\",\"expires_in\":3600}")))) + + ;; GitHub user (= path "/user") (let [token (bearer-token exchange) - user (or (parse-token-user token) {:id 1002 :login "newbie"})] + user (or (parse-token-user token) {:id "1002" :login "newbie" :numeric? true})] (send-json exchange 200 (str "{\"id\":" (:id user) ",\"login\":\"" (:login user) "\"}"))) + ;; Reddit /api/v1/me + (str/ends-with? path "/api/v1/me") + (let [token (bearer-token exchange) + user (or (parse-token-user token) {:id "t2_test" :login "redditor"})] + (send-json exchange 200 + (str "{\"id\":\"" (:id user) "\",\"name\":\"" (:login user) "\"}"))) + :else (send-json exchange 404 "{\"error\":\"not found\"}")))))] (.createContext server "/" handler) diff --git a/test/support/mock_reddit.clj b/test/support/mock_reddit.clj index 5efa92db3e1f79b8f423a2f1adcbda959123c1ad..a630cf0938722193e9af88382d60e777ff371be4 100644 --- a/test/support/mock_reddit.clj +++ b/test/support/mock_reddit.clj @@ -1,14 +1,56 @@ (ns test.support.mock-reddit - "In-process HTTP stub for Reddit API fixtures (`test/fixtures/reddit/`)." + "In-process HTTP stub for Reddit API fixtures + OAuth login endpoints." (:require [clojure.java.io :as io] [clojure.string :as str]) (:import [com.sun.net.httpserver HttpServer HttpHandler HttpExchange] - [java.net InetSocketAddress])) + [java.net InetSocketAddress URLDecoder])) (defn fixtures-dir ([] (fixtures-dir (System/getProperty "user.dir"))) ([root] (str root "/test/fixtures/reddit"))) +(defn- query-param [query key] + (when query + (some (fn [pair] + (let [[k v] (str/split pair "=" 2)] + (when (= k key) + (URLDecoder/decode (or v "") "UTF-8")))) + (str/split query #"&")))) + +(defn- parse-mock-user [raw] + (let [s (or raw "t2_test:redditor") + [id login] (str/split s #":" 2)] + {:id id :login (or login "redditor")})) + +(defn- send-bytes [^HttpExchange ex status ^bytes body content-type] + (.set (.getResponseHeaders ex) "Content-Type" content-type) + (.sendResponseHeaders ex status (alength body)) + (doto (.getResponseBody ex) + (.write body) + (.close))) + +(defn- send-json [^HttpExchange ex status body] + (send-bytes ex status (.getBytes body "UTF-8") "application/json")) + +(defn- send-redirect [^HttpExchange ex location] + (.set (.getResponseHeaders ex) "Location" location) + (.sendResponseHeaders ex 302 -1) + (.close (.getResponseBody ex))) + +(defn- read-form [^HttpExchange ex] + (let [body (slurp (.getInputStream ex))] + {:code (query-param body "code") + :grant (query-param body "grant_type")})) + +(defn- bearer-token [^HttpExchange ex] + (some-> (.getRequestHeaders ex) + (.getFirst "Authorization") + (str/replace #"^[Bb]earer " ""))) + +(defn- parse-token-user [token] + (when (str/starts-with? token "mock:") + (parse-mock-user (subs token 5)))) + (defn start-mock-reddit "Start a mock Reddit API on `port`. Returns a zero-arg `stop` function." ([port] (start-mock-reddit port (fixtures-dir))) @@ -19,15 +61,38 @@ handler (proxy [HttpHandler] [] (handle [^HttpExchange exchange] - ;; `/r//about.json` → subreddit entity; `/r/.json` → listing. - (let [path (.getPath (.getRequestURI exchange)) - body (if (str/includes? path "/about") - about - listing)] - (.sendResponseHeaders exchange 200 (alength body)) - (let [out (.getResponseBody exchange)] - (.write out body) - (.close out)))))] + (let [uri (.getRequestURI exchange) + path (.getPath uri) + query (.getQuery uri) + method (.getRequestMethod exchange)] + (cond + (str/ends-with? path "/api/v1/authorize") + (let [redirect-uri (query-param query "redirect_uri") + state (query-param query "state") + user (parse-mock-user (query-param query "mock_user")) + code (str "mock:" (:id user) ":" (:login user)) + loc (str redirect-uri "?code=" (java.net.URLEncoder/encode code "UTF-8") + "&state=" (java.net.URLEncoder/encode state "UTF-8"))] + (send-redirect exchange loc)) + + (and (= method "POST") (str/ends-with? path "/api/v1/access_token")) + (let [form (read-form exchange) + grant (or (:grant form) "") + code (or (:code form) "mock:t2_test:redditor")] + (if (= grant "client_credentials") + (send-json exchange 200 "{\"access_token\":\"app-token\",\"token_type\":\"bearer\",\"expires_in\":3600}") + (send-json exchange 200 (str "{\"access_token\":\"" code "\",\"token_type\":\"bearer\",\"expires_in\":3600}")))) + + (str/ends-with? path "/api/v1/me") + (let [user (or (parse-token-user (bearer-token exchange)) + {:id "t2_test" :login "redditor"})] + (send-json exchange 200 + (str "{\"id\":\"" (:id user) "\",\"name\":\"" (:login user) "\"}"))) + + ;; `/r//about.json` → subreddit entity; `/r/.json` → listing. + :else + (let [body (if (str/includes? path "/about") about listing)] + (send-bytes exchange 200 body "application/json"))))))] (.createContext server "/" handler) (.setExecutor server nil) (.start server) Side B — contributor: tommy-mor Side B — commit message: [239c074b] url schema stuff Side B — unified diff (full patch): diff --git a/AGENTS.md b/AGENTS.md index 426a88e7c1da54fe0a28c5c76fa4e1f1bc117fcf..e60b9ba6012593361ef10e8fdd9439cd9932e09b 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -58,3 +58,4 @@ Use **tmux** for `cargo run --package sorter2-server` (dev server). Rebuild afte - First `cargo test` / `cargo build --release` is slow; Clojure smoke test always does a release build. - `legacy/` and `ideas/` are not part of the workspace build. +- **ItemId** for web URLs is a canonical full URL (`https://reddit.com/r/rust`). Rules live in [`server/src/url_rules/`](server/src/url_rules/) (composable Rust, not a config DSL). After changing canonicalization rules, rebuild the projection: `cargo run --package sorter2-server -- replay-index`. diff --git a/Cargo.lock b/Cargo.lock index 0dd4fce5fb6400ae153cca4e3dbf5a5158e6d8b4..49a908ef935c430dbe63c6a28d8a24e38b489486 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1951,6 +1951,7 @@ dependencies = [ "tower-http 0.5.2", "tracing", "tracing-subscriber", + "url", "urlencoding", ] diff --git a/REPLAY.sh b/REPLAY.sh new file mode 100755 index 0000000000000000000000000000000000000000..f2dbd8aea60c02d2feef74805f7ef5c2b7022537 --- /dev/null +++ b/REPLAY.sh @@ -0,0 +1,2 @@ +cargo run --package sorter2-server -- replay-index + diff --git a/server/Cargo.toml b/server/Cargo.toml index 27f552c20b97ef28cdde4cb6b1a4980375135111..ad4912791aff59fb1d3293f66ad381ae618cd60b 100644 --- a/server/Cargo.toml +++ b/server/Cargo.toml @@ -24,6 +24,7 @@ async-stream = "0.3" futures-util = { version = "0.3", default-features = false, features = ["std"] } rand = "0.8" urlencoding = "2" +url = "2" durable = { path = "../durable" } [dev-dependencies] diff --git a/server/src/entity_store.rs b/server/src/entity_store.rs index d5d17c3676e4a8ddec998e9f5a9dbafe9c2d9d0e..d29f39aecca6f12cdcf263cf77c3654eb4ee6cfa 100644 --- a/server/src/entity_store.rs +++ b/server/src/entity_store.rs @@ -124,7 +124,7 @@ mod tests { fn round_trip_payload() { let tmp = tempfile::tempdir().unwrap(); let store = EntityStore::open(tmp.path()).unwrap(); - let id = ItemId::parse("reddit.com/r/rust").unwrap(); + let id = ItemId::from_url("https://reddit.com/r/rust").unwrap(); let payload = json!({"kind": "t5", "data": {"display_name": "rust"}}); store.put(&id, &payload).unwrap(); diff --git a/server/src/event_log.rs b/server/src/event_log.rs index 36f5b406084065b608735987cdb483c236e03081..2c9290b6fdbf2c2ad1c0f1ffd7374b2d9cc97f36 100644 --- a/server/src/event_log.rs +++ b/server/src/event_log.rs @@ -199,7 +199,7 @@ mod tests { log.append(&sample_record( 1, Event::NodeEnsured { - id: "reddit.com/r/rust".into(), + id: "https://reddit.com/r/rust".into(), }, )) .await @@ -237,7 +237,7 @@ mod tests { let path = tmp.path().join("events.jsonl"); let log = EventLog::new(&path); let event = Event::NodeEnsured { - id: "reddit.com/r/rust".into(), + id: "https://reddit.com/r/rust".into(), }; log.append(&sample_record(1, event)).await.unwrap(); @@ -255,7 +255,7 @@ mod tests { let path = tmp.path().join("events.jsonl"); std::fs::write( &path, - r#"{"type":"node_ensured","id":"reddit.com/r/rust"} + r#"{"type":"node_ensured","id":"https://reddit.com/r/rust"} {"schema":1,"seq":1,"ts":1,"event":{"type":"vote_recorded","ts":1,"a":"a","b":"b","ratio_left":2,"ratio_right":1,"scope":""}} "#, ) @@ -295,7 +295,7 @@ mod tests { log.append(&sample_record( 1, Event::NodeEnsured { - id: "reddit.com/r/rust".into(), + id: "https://reddit.com/r/rust".into(), }, )) .await @@ -303,7 +303,7 @@ mod tests { log.append(&sample_record( 3, Event::NodeEnsured { - id: "reddit.com/r/python".into(), + id: "https://reddit.com/r/python".into(), }, )) .await diff --git a/server/src/journal.rs b/server/src/journal.rs index 521a108019de1ea870d14c4fafbfe572c20ce0de..50bc89f976edb82b7b0e49e954a8eccbbe82bf87 100644 --- a/server/src/journal.rs +++ b/server/src/journal.rs @@ -141,10 +141,10 @@ mod tests { let j2 = journal.clone(); let (r1, r2) = tokio::join!( j1.append(Event::NodeEnsured { - id: "reddit.com/r/rust".into(), + id: "https://reddit.com/r/rust".into(), }), j2.append(Event::NodeEnsured { - id: "reddit.com/r/python".into(), + id: "https://reddit.com/r/python".into(), }), ); r1.unwrap(); @@ -153,10 +153,10 @@ mod tests { assert_eq!(projection_store.last_applied_event_count().unwrap(), 2); let tree = projection_store.load_tree().unwrap(); assert!(tree - .get(&ItemId::parse("reddit.com/r/rust").unwrap()) + .get(&ItemId::parse("https://reddit.com/r/rust").unwrap()) .is_some()); assert!(tree - .get(&ItemId::parse("reddit.com/r/python").unwrap()) + .get(&ItemId::parse("https://reddit.com/r/python").unwrap()) .is_some()); } @@ -170,7 +170,7 @@ mod tests { 1, 1, Event::NodeEnsured { - id: "reddit.com/r/rust".into(), + id: "https://reddit.com/r/rust".into(), }, )) .await @@ -186,7 +186,7 @@ mod tests { 1, 1, Event::NodeEnsured { - id: "reddit.com/r/rust".into(), + id: "https://reddit.com/r/rust".into(), }, )], ) @@ -202,7 +202,7 @@ mod tests { ); journal .append(Event::NodeEnsured { - id: "reddit.com/r/python".into(), + id: "https://reddit.com/r/python".into(), }) .await .unwrap(); @@ -227,13 +227,13 @@ mod tests { journal .append_many(vec![ Event::NodeEnsured { - id: "reddit.com/r/rust".into(), + id: "https://reddit.com/r/rust".into(), }, Event::NodeEnsured { - id: "reddit.com/r/python".into(), + id: "https://reddit.com/r/python".into(), }, Event::NodeEnsured { - id: "reddit.com/r/clojure".into(), + id: "https://reddit.com/r/clojure".into(), }, ]) .await @@ -245,7 +245,7 @@ mod tests { assert_eq!(projection_store.last_applied_event_count().unwrap(), 3); let tree = projection_store.load_tree().unwrap(); assert!(tree - .get(&ItemId::parse("reddit.com/r/clojure").unwrap()) + .get(&ItemId::parse("https://reddit.com/r/clojure").unwrap()) .is_some()); } } diff --git a/server/src/lib.rs b/server/src/lib.rs index 9bd5f76fd1406b9b1be4c272f4ba8647edde2678..5c02c8e704e4664453bad75d819df8a067668176 100644 --- a/server/src/lib.rs +++ b/server/src/lib.rs @@ -9,6 +9,7 @@ pub mod journal; pub mod pair; pub mod parser; pub mod path_types; +pub mod url_rules; pub mod projection_apply; pub mod projection_store; pub mod ranking; diff --git a/server/src/pair.rs b/server/src/pair.rs index 43f780ba6ea6ce1cdc2e1f4cbb252ba8a10684b9..815a97b80e3e9f348e0937a4f147f2862018edb0 100644 --- a/server/src/pair.rs +++ b/server/src/pair.rs @@ -381,42 +381,42 @@ mod tests { #[test] fn suggest_prefers_unvoted_pair() { - let parent = ItemId::parse("reddit.com/r/rust").unwrap(); + let parent = ItemId::parse("https://reddit.com/r/rust").unwrap(); let mut tree = seed_children( &parent, &[ - "reddit.com/r/rust/a", - "reddit.com/r/rust/b", - "reddit.com/r/rust/c", + "https://reddit.com/r/rust/a", + "https://reddit.com/r/rust/b", + "https://reddit.com/r/rust/c", ], ); let vote = - VoteData::from_recorded(1, "reddit.com/r/rust/a", "reddit.com/r/rust/b", 2, 1).unwrap(); + VoteData::from_recorded(1, "https://reddit.com/r/rust/a", "https://reddit.com/r/rust/b", 2, 1).unwrap(); tree.apply_vote(&parent, vote); let group = tree.get(&parent).unwrap().local_ranking.clone(); let pool = children_of(&tree, &parent); let (l, r) = suggest_next_pair_in_pool(&group, &pool, None).unwrap(); - let voted_ab = (l.as_str() == "reddit.com/r/rust/a" && r.as_str() == "reddit.com/r/rust/b") - || (l.as_str() == "reddit.com/r/rust/b" && r.as_str() == "reddit.com/r/rust/a"); + let voted_ab = (l.as_str() == "https://reddit.com/r/rust/a" && r.as_str() == "https://reddit.com/r/rust/b") + || (l.as_str() == "https://reddit.com/r/rust/b" && r.as_str() == "https://reddit.com/r/rust/a"); assert!(!voted_ab); } #[test] fn suggest_bridges_separate_components() { - let parent = ItemId::parse("reddit.com/r/rust").unwrap(); + let parent = ItemId::parse("https://reddit.com/r/rust").unwrap(); let mut tree = seed_children( &parent, &[ - "reddit.com/r/rust/a", - "reddit.com/r/rust/b", - "reddit.com/r/rust/c", - "reddit.com/r/rust/d", + "https://reddit.com/r/rust/a", + "https://reddit.com/r/rust/b", + "https://reddit.com/r/rust/c", + "https://reddit.com/r/rust/d", ], ); let ab = - VoteData::from_recorded(1, "reddit.com/r/rust/a", "reddit.com/r/rust/b", 2, 1).unwrap(); + VoteData::from_recorded(1, "https://reddit.com/r/rust/a", "https://reddit.com/r/rust/b", 2, 1).unwrap(); let cd = - VoteData::from_recorded(2, "reddit.com/r/rust/c", "reddit.com/r/rust/d", 2, 1).unwrap(); + VoteData::from_recorded(2, "https://reddit.com/r/rust/c", "https://reddit.com/r/rust/d", 2, 1).unwrap(); tree.apply_vote(&parent, ab); tree.apply_vote(&parent, cd); let group = tree.get(&parent).unwrap().local_ranking.clone(); @@ -424,37 +424,37 @@ mod tests { let pair = suggest_next_pair_in_pool(&group, &pool, None).unwrap(); let chosen = pair_set(&pair); let from_ab = - chosen.contains("reddit.com/r/rust/a") || chosen.contains("reddit.com/r/rust/b"); + chosen.contains("https://reddit.com/r/rust/a") || chosen.contains("https://reddit.com/r/rust/b"); let from_cd = - chosen.contains("reddit.com/r/rust/c") || chosen.contains("reddit.com/r/rust/d"); + chosen.contains("https://reddit.com/r/rust/c") || chosen.contains("https://reddit.com/r/rust/d"); assert!(from_ab && from_cd, "expected bridge pair, got {:?}", chosen); } #[test] fn suggest_prefers_attach_over_isolate_pair_among_many_unranked() { - let parent = ItemId::parse("reddit.com/r/rust").unwrap(); + let parent = ItemId::parse("https://reddit.com/r/rust").unwrap(); let mut tree = seed_children( &parent, &[ - "reddit.com/r/rust/a", - "reddit.com/r/rust/b", - "reddit.com/r/rust/c", - "reddit.com/r/rust/d", - "reddit.com/r/rust/e", + "https://reddit.com/r/rust/a", + "https://reddit.com/r/rust/b", + "https://reddit.com/r/rust/c", + "https://reddit.com/r/rust/d", + "https://reddit.com/r/rust/e", ], ); let ab = - VoteData::from_recorded(1, "reddit.com/r/rust/a", "reddit.com/r/rust/b", 2, 1).unwrap(); + VoteData::from_recorded(1, "https://reddit.com/r/rust/a", "https://reddit.com/r/rust/b", 2, 1).unwrap(); tree.apply_vote(&parent, ab); let group = tree.get(&parent).unwrap().local_ranking.clone(); let pool = children_of(&tree, &parent); let pair = suggest_next_pair_in_pool(&group, &pool, None).unwrap(); let chosen = pair_set(&pair); let from_ab = - chosen.contains("reddit.com/r/rust/a") || chosen.contains("reddit.com/r/rust/b"); - let from_cde = chosen.contains("reddit.com/r/rust/c") - || chosen.contains("reddit.com/r/rust/d") - || chosen.contains("reddit.com/r/rust/e"); + chosen.contains("https://reddit.com/r/rust/a") || chosen.contains("https://reddit.com/r/rust/b"); + let from_cde = chosen.contains("https://reddit.com/r/rust/c") + || chosen.contains("https://reddit.com/r/rust/d") + || chosen.contains("https://reddit.com/r/rust/e"); assert!( from_ab && from_cde, "expected ranked+unranked attach, got {:?}", @@ -464,40 +464,40 @@ mod tests { #[test] fn suggest_connects_isolate_to_existing_component() { - let parent = ItemId::parse("reddit.com/r/rust").unwrap(); + let parent = ItemId::parse("https://reddit.com/r/rust").unwrap(); let mut tree = seed_children( &parent, &[ - "reddit.com/r/rust/a", - "reddit.com/r/rust/b", - "reddit.com/r/rust/c", + "https://reddit.com/r/rust/a", + "https://reddit.com/r/rust/b", + "https://reddit.com/r/rust/c", ], ); let ab = - VoteData::from_recorded(1, "reddit.com/r/rust/a", "reddit.com/r/rust/b", 2, 1).unwrap(); + VoteData::from_recorded(1, "https://reddit.com/r/rust/a", "https://reddit.com/r/rust/b", 2, 1).unwrap(); tree.apply_vote(&parent, ab); let group = tree.get(&parent).unwrap().local_ranking.clone(); let pool = children_of(&tree, &parent); let pair = suggest_next_pair_in_pool(&group, &pool, None).unwrap(); let chosen = pair_set(&pair); - assert!(chosen.contains("reddit.com/r/rust/c")); - assert!(chosen.contains("reddit.com/r/rust/a") || chosen.contains("reddit.com/r/rust/b")); + assert!(chosen.contains("https://reddit.com/r/rust/c")); + assert!(chosen.contains("https://reddit.com/r/rust/a") || chosen.contains("https://reddit.com/r/rust/b")); } #[test] fn suggest_zips_adjacent_ranks_when_tree_complete() { - let parent = ItemId::parse("reddit.com/r/rust").unwrap(); + let parent = ItemId::parse("https://reddit.com/r/rust").unwrap(); let mut tree = seed_children( &parent, &[ - "reddit.com/r/rust/a", - "reddit.com/r/rust/b", - "reddit.com/r/rust/c", + "https://reddit.com/r/rust/a", + "https://reddit.com/r/rust/b", + "https://reddit.com/r/rust/c", ], ); for (a, b, l, r) in [ - ("reddit.com/r/rust/a", "reddit.com/r/rust/b", 3, 1), - ("reddit.com/r/rust/a", "reddit.com/r/rust/c", 2, 1), + ("https://reddit.com/r/rust/a", "https://reddit.com/r/rust/b", 3, 1), + ("https://reddit.com/r/rust/a", "https://reddit.com/r/rust/c", 2, 1), ] { let v = VoteData::from_recorded(1, a, b, l, r).unwrap(); tree.apply_vote(&parent, v); @@ -506,26 +506,26 @@ mod tests { let pool = children_of(&tree, &parent); let pair = suggest_next_pair_in_pool(&group, &pool, None).unwrap(); let chosen = pair_set(&pair); - assert!(chosen.contains("reddit.com/r/rust/b")); - assert!(chosen.contains("reddit.com/r/rust/c")); + assert!(chosen.contains("https://reddit.com/r/rust/b")); + assert!(chosen.contains("https://reddit.com/r/rust/c")); } #[test] fn suggest_zip_prefers_1v2_before_2v3_when_both_unvoted() { - let parent = ItemId::parse("reddit.com/r/rust").unwrap(); + let parent = ItemId::parse("https://reddit.com/r/rust").unwrap(); let mut tree = seed_children( &parent, &[ - "reddit.com/r/rust/a", - "reddit.com/r/rust/b", - "reddit.com/r/rust/c", - "reddit.com/r/rust/d", + "https://reddit.com/r/rust/a", + "https://reddit.com/r/rust/b", + "https://reddit.com/r/rust/c", + "https://reddit.com/r/rust/d", ], ); for (a, b, l, r) in [ - ("reddit.com/r/rust/c", "reddit.com/r/rust/d", 3, 1), - ("reddit.com/r/rust/b", "reddit.com/r/rust/c", 2, 1), - ("reddit.com/r/rust/a", "reddit.com/r/rust/c", 2, 1), + ("https://reddit.com/r/rust/c", "https://reddit.com/r/rust/d", 3, 1), + ("https://reddit.com/r/rust/b", "https://reddit.com/r/rust/c", 2, 1), + ("https://reddit.com/r/rust/a", "https://reddit.com/r/rust/c", 2, 1), ] { let v = VoteData::from_recorded(1, a, b, l, r).unwrap(); tree.apply_vote(&parent, v); @@ -534,16 +534,16 @@ mod tests { let pool = children_of(&tree, &parent); let pair = suggest_next_pair_in_pool(&group, &pool, None).unwrap(); let chosen = pair_set(&pair); - assert!(chosen.contains("reddit.com/r/rust/a")); - assert!(chosen.contains("reddit.com/r/rust/b")); + assert!(chosen.contains("https://reddit.com/r/rust/a")); + assert!(chosen.contains("https://reddit.com/r/rust/b")); } #[test] fn resolve_pair_picks_from_pool() { - let parent = ItemId::parse("reddit.com/r/rust").unwrap(); - let tree = seed_children(&parent, &["reddit.com/r/rust/a", "reddit.com/r/rust/b"]); + let parent = ItemId::parse("https://reddit.com/r/rust").unwrap(); + let tree = seed_children(&parent, &["https://reddit.com/r/rust/a", "https://reddit.com/r/rust/b"]); let pair = resolve_pair(&tree, &parent, None, None).unwrap(); - let pool: HashSet<_> = ["reddit.com/r/rust/a", "reddit.com/r/rust/b"] + let pool: HashSet<_> = ["https://reddit.com/r/rust/a", "https://reddit.com/r/rust/b"] .into_iter() .collect(); assert!(pool.contains(pair.0.as_str())); diff --git a/server/src/parser.rs b/server/src/parser.rs index 9df2dcc9313f7fe250ce3b6aa167f6ec5d57951f..b2a963dd6cab415576c8d3a9a241966758565bb8 100644 --- a/server/src/parser.rs +++ b/server/src/parser.rs @@ -23,7 +23,7 @@ mod tests { fn parses_short_path() { assert_eq!( parse_reddit_url("r/rust").unwrap().as_str(), - "reddit.com/r/rust" + "https://reddit.com/r/rust" ); } @@ -33,7 +33,7 @@ mod tests { parse_reddit_url("https://www.reddit.com/r/programming/hot") .unwrap() .as_str(), - "reddit.com/r/programming" + "https://reddit.com/r/programming" ); } @@ -43,7 +43,10 @@ mod tests { "https://old.reddit.com/r/AmItheAsshole/comments/1trnvdl/aita_for_cancelling/", ) .unwrap(); - assert_eq!(id.as_str(), "reddit.com/r/amitheasshole/comments/1trnvdl"); + assert_eq!( + id.as_str(), + "https://reddit.com/r/amitheasshole/comments/1trnvdl" + ); } #[test] diff --git a/server/src/path_types.rs b/server/src/path_types.rs index fafd924452f6fd85e7a5b27ed2653a19581e6e56..71ffc01f35c5589686ff05dae3b5610fa01f30ce 100644 --- a/server/src/path_types.rs +++ b/server/src/path_types.rs @@ -1,13 +1,14 @@ use serde::{Deserialize, Serialize}; use std::fmt; -/// Canonical hierarchical identity for any URL/path in the fractal tree. +use crate::url_rules::{looks_like_url, navigable_breadcrumbs, parent_url, resolve_canonical}; + +/// Canonical identity: a real URL (with scheme) or an opaque non-URL key. #[derive(Debug, Clone, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize, Default)] pub struct ItemId(String); impl ItemId { - /// Parse an already-canonical path (no URL normalization). Empty string is invalid here; - /// use [`Self::root`] for the tree root. + /// Parse an already-canonical id (no normalization). Empty string is invalid; use [`Self::root`]. pub fn parse(s: &str) -> Option { let t = s.trim(); if t.is_empty() { @@ -16,12 +17,12 @@ impl ItemId { Some(Self(t.to_string())) } - /// Build an opaque item key (legacy demo votes, non-URL items). + /// Build an opaque item key (demo votes, non-URL items). pub fn opaque(s: impl Into) -> Self { Self(s.into()) } - /// Root of the internet tree (empty path). + /// Root of the internet tree. pub fn root() -> Self { Self(String::new()) } @@ -34,23 +35,18 @@ impl ItemId { &self.0 } - /// Creates a canonical ID from a raw URL or path. Normalizes domains and - /// trims tracking query params. + /// Canonical URL from a raw pasted or fetched URL. pub fn from_url(raw_url: &str) -> Option { - Self::canonicalize(raw_url).map(Self) + resolve_canonical(raw_url).map(Self) } - /// Normalize strings from forms, events, and Reddit imports into the same - /// stored id shape (e.g. drop post title slug after comment id). + /// Normalize strings from forms, events, and imports into canonical identity. pub fn from_storage(s: &str) -> Option { let t = s.trim(); if t.is_empty() { return None; } - if t.contains("://") || t.starts_with("r/") { - return Self::from_url(t).or_else(|| Self::parse(t)); - } - if t.starts_with("reddit.com/") && t.contains("/comments/") { + if looks_like_url(t) { return Self::from_url(t).or_else(|| Self::parse(t)); } Self::parse(t).or_else(|| Self::from_url(t)) @@ -62,35 +58,52 @@ impl ItemId { if s.is_empty() { return Self::root(); } - Self(format!("reddit.com/r/{s}")) + if looks_like_url(s) || s.contains('/') { + Self::from_storage(s).unwrap_or_else(|| Self::opaque(s)) + } else { + Self(format!("https://reddit.com/r/{s}")) + } } - /// Extract the parent, e.g. `reddit.com/r/aww/comments/1trnvdl` → - /// `reddit.com/r/aww`. + /// Immediate parent scope in the tree. pub fn parent(&self) -> Option { - if self.0.is_empty() { + if self.is_root() { return None; } - + if looks_like_url(self.0.as_str()) { + return parent_url(self.0.as_str()).map(Self); + } let parts: Vec<&str> = self.0.trim_end_matches('/').split('/').collect(); if parts.len() <= 1 { return None; } - - if self.0.contains("/comments/") { - return Some(Self(parts[..parts.len().saturating_sub(2)].join("/"))); - } - Some(Self(parts[..parts.len() - 1].join("/"))) } pub fn segments(&self) -> Vec<&str> { + if self.is_root() { + return vec![]; + } + if let Some(rest) = self.0.strip_prefix("https://") { + return rest.split('/').filter(|s| !s.is_empty()).collect(); + } + if let Some(rest) = self.0.strip_prefix("http://") { + return rest.split('/').filter(|s| !s.is_empty()).collect(); + } self.0.split('/').filter(|s| !s.is_empty()).collect() } - /// Cumulative paths for breadcrumb rendering, e.g. - /// `reddit.com/r/movies` → `["reddit.com", "reddit.com/r", "reddit.com/r/movies"]`. + /// Cumulative navigable paths for breadcrumbs and tree wiring (includes self). pub fn breadcrumb_paths(&self) -> Vec { + if self.is_root() { + return vec![]; + } + if looks_like_url(self.0.as_str()) { + return navigable_breadcrumbs(self.0.as_str()) + .into_iter() + .map(ItemId) + .collect(); + } let segs = self.segments(); let mut paths = Vec::with_capacity(segs.len()); let mut current = String::new(); @@ -111,13 +124,13 @@ impl ItemId { if self.is_root() { return String::new(); } - if self.as_str().contains("://") { - return self.as_str().to_string(); + if self.0.contains("://") { + return self.0.clone(); } if self.segments().first().is_some_and(|s| s.contains('.')) { - format!("https://{}", self.as_str()) + format!("https://{}", self.0) } else { - self.as_str().to_string() + self.0.clone() } } @@ -144,70 +157,6 @@ impl ItemId { pub fn from_browse_uri(path: &str) -> Option { path.strip_prefix("/~/").map(ItemId::from_browse_tail) } - - fn canonicalize(raw: &str) -> Option { - let s = raw.trim(); - if s.is_empty() { - return None; - } - - let owned = if let Some(rest) = s.strip_prefix("r/") { - format!("reddit.com/r/{rest}") - } else if let Some(rest) = s.strip_prefix("/r/") { - format!("reddit.com/r/{rest}") - } else { - s.to_string() - }; - - let (host_path, _query) = split_query(&owned); - let host_path = host_path.trim_end_matches('/'); - - let path = if host_path.contains("://") { - parse_url_host_path(host_path)? - } else if host_path.starts_with("reddit.com") || host_path.starts_with("www.reddit.com") { - normalize_reddit_host_path(host_path) - } else if host_path.contains('/') { - host_path.to_string() - } else { - return None; - }; - - Some(normalize_reddit_path(&path)) - } -} - -fn split_query(s: &str) -> (&str, Option<&str>) { - if let Some((path, q)) = s.split_once('?') { - (path, Some(q)) - } else { - (s, None) - } -} - -fn parse_url_host_path(url: &str) -> Option { - let rest = url - .strip_prefix("https://") - .or_else(|| url.strip_prefix("http://")) - .unwrap_or(url); - let (host, path) = rest.split_once('/').unwrap_or((rest, "")); - let host = normalize_host(host); - if path.is_empty() { - Some(host) - } else { - Some(format!("{host}/{path}")) - } -} - -fn normalize_host(host: &str) -> String { - let h = host - .strip_prefix("www.") - .unwrap_or(host) - .to_ascii_lowercase(); - if h == "old.reddit.com" || h == "new.reddit.com" || h == "reddit.com" { - "reddit.com".to_string() - } else { - h - } } fn normalize_browse_tail(tail: &str) -> String { @@ -215,7 +164,6 @@ fn normalize_browse_tail(tail: &str) -> String { if t.is_empty() { return String::new(); } - // Some HTTP stacks collapse `https://` → `https:/` inside a path segment. if t.starts_with("https:/") && !t.starts_with("https://") { return format!("https://{}", &t[7..]); } @@ -225,33 +173,6 @@ fn normalize_browse_tail(tail: &str) -> String { t.to_string() } -fn normalize_reddit_host_path(s: &str) -> String { - let (host, path) = s.split_once('/').unwrap_or((s, "")); - let host = normalize_host(host); - if path.is_empty() { - host - } else { - format!("{host}/{path}") - } -} - -/// Lowercase subreddit segment, drop listing suffixes, drop title slug after post id. -fn normalize_reddit_path(path: &str) -> String { - let mut parts: Vec = path.split('/').map(str::to_string).collect(); - if parts.len() >= 3 && parts[1] == "r" { - parts[2] = parts[2].to_ascii_lowercase(); - } - if let Some(i) = parts.iter().position(|p| p == "comments") { - if parts.len() > i + 2 { - parts.truncate(i + 2); - } - } else if parts.len() > 3 && parts.get(1).map(|s| s.as_str()) == Some("r") { - // reddit.com/r/{sub}/hot → reddit.com/r/{sub} - parts.truncate(3); - } - parts.join("/") -} - impl fmt::Display for ItemId { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { f.write_str(&self.0) @@ -268,43 +189,59 @@ mod tests { "https://old.reddit.com/r/AmItheAsshole/comments/1trnvdl/aita_for_cancelling/", ) .unwrap(); - assert_eq!(id.as_str(), "reddit.com/r/amitheasshole/comments/1trnvdl"); + assert_eq!( + id.as_str(), + "https://reddit.com/r/amitheasshole/comments/1trnvdl" + ); } #[test] fn from_url_strips_query() { let id = ItemId::from_url("https://www.reddit.com/r/rust/?sort=top").unwrap(); - assert_eq!(id.as_str(), "reddit.com/r/rust"); + assert_eq!(id.as_str(), "https://reddit.com/r/rust"); } #[test] fn from_url_short_path() { assert_eq!( ItemId::from_url("r/rust").unwrap().as_str(), - "reddit.com/r/rust" + "https://reddit.com/r/rust" ); } #[test] fn parent_of_post_is_subreddit() { - let id = ItemId::parse("reddit.com/r/aww/comments/1trnvdl").unwrap(); - assert_eq!(id.parent().unwrap().as_str(), "reddit.com/r/aww"); + let id = ItemId::from_url("https://reddit.com/r/aww/comments/1trnvdl").unwrap(); + assert_eq!(id.parent().unwrap().as_str(), "https://reddit.com/r/aww"); } #[test] fn parent_of_subreddit_is_r_segment() { - let id = ItemId::parse("reddit.com/r/movies").unwrap(); - assert_eq!(id.parent().unwrap().as_str(), "reddit.com/r"); + let id = ItemId::from_url("https://reddit.com/r/movies").unwrap(); + assert_eq!(id.parent().unwrap().as_str(), "https://reddit.com/r"); + } + + #[test] + fn breadcrumb_paths_skip_phantom_comments() { + let id = ItemId::from_url("https://reddit.com/r/aww/comments/1trnvdl").unwrap(); + let crumbs = id.breadcrumb_paths(); + let paths: Vec<_> = crumbs.iter().map(|p| p.as_str()).collect(); + assert!(!paths.iter().any(|p| p.ends_with("/comments"))); + assert!(paths.contains(&"https://reddit.com/r/aww")); } #[test] - fn breadcrumb_paths() { - let id = ItemId::parse("reddit.com/r/movies").unwrap(); + fn breadcrumb_paths_subreddit() { + let id = ItemId::from_url("https://reddit.com/r/movies").unwrap(); let crumbs = id.breadcrumb_paths(); let paths: Vec<_> = crumbs.iter().map(|p| p.as_str()).collect(); assert_eq!( paths, - vec!["reddit.com", "reddit.com/r", "reddit.com/r/movies"] + vec![ + "https://reddit.com", + "https://reddit.com/r", + "https://reddit.com/r/movies" + ] ); } @@ -312,33 +249,33 @@ mod tests { fn legacy_scope_maps_to_reddit_sub() { assert_eq!( ItemId::from_legacy_scope("rust").as_str(), - "reddit.com/r/rust" + "https://reddit.com/r/rust" ); assert!(ItemId::from_legacy_scope("").is_root()); } #[test] - fn browse_href_wraps_canonical_path() { - let id = ItemId::parse("reddit.com/r/rust").unwrap(); + fn browse_href_wraps_canonical_url() { + let id = ItemId::from_url("https://reddit.com/r/rust").unwrap(); assert_eq!(id.browse_href(), "/~/https://reddit.com/r/rust"); } #[test] fn from_browse_tail_parses_full_url() { let id = ItemId::from_browse_tail("https://reddit.com/r/AmITheAsshole"); - assert_eq!(id.as_str(), "reddit.com/r/amitheasshole"); + assert_eq!(id.as_str(), "https://reddit.com/r/amitheasshole"); } #[test] fn from_storage_strips_post_title_slug() { let id = ItemId::from_storage("reddit.com/r/rust/comments/aaa/announcing_rust_199").unwrap(); - assert_eq!(id.as_str(), "reddit.com/r/rust/comments/aaa"); + assert_eq!(id.as_str(), "https://reddit.com/r/rust/comments/aaa"); } #[test] fn from_browse_uri_strips_prefix() { let id = ItemId::from_browse_uri("/~/https://reddit.com/r/rust").unwrap(); - assert_eq!(id.as_str(), "reddit.com/r/rust"); + assert_eq!(id.as_str(), "https://reddit.com/r/rust"); } } diff --git a/server/src/projection_apply.rs b/server/src/projection_apply.rs index ebe122e417bda1d9369a53443de93d28213c5a0d..5644557a41b3e9497c7421b444155ae629fa79f1 100644 --- a/server/src/projection_apply.rs +++ b/server/src/projection_apply.rs @@ -19,13 +19,21 @@ use crate::{ storage_schema::{ensure_path_writes, entity_view_writes, vote_writes}, }; -/// Legacy-compatible scope parsing for persisted vote events. +fn parse_event_id(id: &str) -> Result { + ItemId::from_storage(id) + .or_else(|| ItemId::parse(id)) + .ok_or_else(|| EventLogError::Apply(format!("invalid id: {id}"))) +} + +/// Scope key from a vote event (canonicalized at apply time). fn parent_from_event_scope(scope: &str) -> ItemId { - if scope.contains('/') { - ItemId::parse(scope).unwrap_or_else(|| ItemId::from_legacy_scope(scope)) - } else { - ItemId::from_legacy_scope(scope) + let s = scope.trim(); + if s.is_empty() { + return ItemId::root(); } + ItemId::from_storage(s) + .or_else(|| ItemId::parse(s)) + .unwrap_or_else(|| ItemId::from_legacy_scope(s)) } pub fn apply_records( @@ -68,15 +76,11 @@ pub fn apply_records( vote_parents.insert(parent); } Event::NodeEnsured { id } => { - let parsed = ItemId::parse(id) - .or_else(|| ItemId::from_url(id)) - .ok_or_else(|| EventLogError::Apply(format!("invalid node id: {id}")))?; + let parsed = parse_event_id(id)?; ensure_path_writes(&mut batch, &parsed); } Event::EntityImported { id, payload, .. } => { - let parsed = ItemId::parse(id) - .or_else(|| ItemId::from_url(id)) - .ok_or_else(|| EventLogError::Apply(format!("invalid entity id: {id}")))?; + let parsed = parse_event_id(id)?; let view = entity_view_from_payload(&parsed, payload); entity_view_writes(&mut batch, &parsed, view.as_ref()); entity_store @@ -92,7 +96,6 @@ pub fn apply_records( .commit_with(durable::Durability::DisableWal) .map_err(|e| EventLogError::Apply(e.to_string()))?; - // Cap recent-vote windows (idempotent, blind; not part of the cursor batch). for parent in vote_parents { projection_store .trim_recent_votes(&parent) diff --git a/server/src/reddit.rs b/server/src/reddit.rs index 72caf7c33b9dd44ea15b62b59e91227cf96a3431..20b7f9e3f8be39268a1767d09f5cf81eaa6ae0df 100644 --- a/server/src/reddit.rs +++ b/server/src/reddit.rs @@ -183,7 +183,7 @@ pub fn entity_view_from_payload( id: &ItemId, payload: &Value, ) -> Option { - if id.as_str().starts_with("reddit.com") { + if id.as_str().contains("reddit.com") { return parse_reddit_view(id, payload); } None @@ -495,41 +495,59 @@ fn rate_limit_reset_secs(resp: &reqwest::Response) -> u64 { .unwrap_or(5) } +fn reddit_path_segments(id: &ItemId) -> Option> { + let s = id.as_str(); + let rest = s + .strip_prefix("https://reddit.com/") + .or_else(|| s.strip_prefix("http://reddit.com/")) + .or_else(|| s.strip_prefix("reddit.com/"))?; + let segments: Vec = rest + .split('/') + .filter(|p| !p.is_empty()) + .map(str::to_string) + .collect(); + Some(segments) +} + pub fn map_item_to_reddit_api(id: &ItemId, api_base: &str) -> String { - let path = id.as_str(); - if !path.starts_with("reddit.com/") && path != "reddit.com" { - return String::new(); - } + let segments = match reddit_path_segments(id) { + Some(s) => s, + None if matches!( + id.as_str(), + "https://reddit.com" | "http://reddit.com" | "reddit.com" + ) => + { + return String::new(); + } + None => return String::new(), + }; let base = api_base.trim_end_matches('/'); - let segments: Vec<&str> = path.split('/').collect(); - - if let Some(i) = segments.iter().position(|&p| p == "comments") { + if let Some(i) = segments.iter().position(|p| p == "comments") { if segments.len() > i + 1 { - let api_path = segments[1..=i + 1].join("/"); + let api_path = segments[..=i + 1].join("/"); return format!("{base}/{api_path}.json?raw_json=1"); } } - if segments.len() == 3 && segments[1] == "r" { - return format!("{base}/r/{}/about.json?raw_json=1", segments[2]); + if segments.len() == 2 && segments[0] == "r" { + return format!("{base}/r/{}/about.json?raw_json=1", segments[1]); } String::new() } /// Listing URL for a node's children. Currently only subreddits -/// (`reddit.com/r/` → `/r/.json`) expose a child listing. +/// (`https://reddit.com/r/` → `/r/.json`) expose a child listing. pub fn map_children_url(id: &ItemId, api_base: &str) -> String { - let path = id.as_str(); - if !path.starts_with("reddit.com/") { - return String::new(); - } + let segments = match reddit_path_segments(id) { + Some(s) => s, + None => return String::new(), + }; let base = api_base.trim_end_matches('/'); - let segments: Vec<&str> = path.split('/').collect(); - if segments.len() == 3 && segments[1] == "r" { - return format!("{base}/r/{}.json?raw_json=1&limit=25", segments[2]); + if segments.len() == 2 && segments[0] == "r" { + return format!("{base}/r/{}.json?raw_json=1&limit=25", segments[1]); } String::new() } @@ -548,8 +566,8 @@ fn parse_children(_parent: &ItemId, payload: &Value) -> Vec<(ItemId, Value)> { Some(p) if !p.is_empty() => p, _ => continue, }; - let path = format!("reddit.com{}", permalink.trim_end_matches('/')); - if let Some(id) = ItemId::from_storage(&path) { + let raw = format!("https://reddit.com{}", permalink.trim_end_matches('/')); + if let Some(id) = ItemId::from_url(&raw) { out.push((id, child.clone())); } } @@ -683,7 +701,7 @@ mod tests { #[test] fn map_subreddit_about_url() { - let id = ItemId::parse("reddit.com/r/rust").unwrap(); + let id = ItemId::from_url("https://reddit.com/r/rust").unwrap(); assert_eq!( map_item_to_reddit_api(&id, "https://www.reddit.com"), "https://www.reddit.com/r/rust/about.json?raw_json=1" @@ -699,7 +717,8 @@ mod tests { let json = include_str!("../../test/fixtures/reddit/r_rust_about.json"); let v: Value = serde_json::from_str(json).unwrap(); let entity = - entity_view_from_payload(&ItemId::parse("reddit.com/r/rust").unwrap(), &v).unwrap(); + entity_view_from_payload(&ItemId::from_url("https://reddit.com/r/rust").unwrap(), &v) + .unwrap(); assert_eq!(entity.title, "The Rust Programming Language"); } @@ -707,7 +726,8 @@ mod tests { fn parse_post_listing_extracts_thumb_and_full_preview() { let json = include_str!("../../test/fixtures/reddit/post_preview.json"); let v: Value = serde_json::from_str(json).unwrap(); - let id = ItemId::parse("reddit.com/r/nsfw/comments/1tpy6a1/angel_eyes").unwrap(); + let id = + ItemId::from_url("https://reddit.com/r/nsfw/comments/1tpy6a1/angel_eyes").unwrap(); let entity = entity_view_from_payload(&id, &v).unwrap(); assert_eq!(entity.title, "Angel Eyes"); assert!(entity.thumb_url.as_ref().unwrap().contains("width=140")); diff --git a/server/src/reducer.rs b/server/src/reducer.rs index 6578f64a41726845517cdbf59a359c69e0aa56db..5179ddeca7a7cb0fb92dfc4aa9d5a80bd9125611 100644 --- a/server/src/reducer.rs +++ b/server/src/reducer.rs @@ -248,16 +248,18 @@ mod from_recorded_tests { #[test] fn ensure_path_wires_children() { let mut tree = GlobalTree::new(); - let id = ItemId::parse("reddit.com/r/rust").unwrap(); + let id = ItemId::from_url("https://reddit.com/r/rust").unwrap(); tree.ensure_path(&id); let root = tree.get(&ItemId::root()).unwrap(); assert!(root .children - .contains(&ItemId::parse("reddit.com").unwrap())); - let reddit = tree.get(&ItemId::parse("reddit.com").unwrap()).unwrap(); + .contains(&ItemId::from_url("https://reddit.com").unwrap())); + let reddit = tree + .get(&ItemId::from_url("https://reddit.com").unwrap()) + .unwrap(); assert!(reddit .children - .contains(&ItemId::parse("reddit.com/r").unwrap())); + .contains(&ItemId::from_url("https://reddit.com/r").unwrap())); let sub = tree.get(&id).unwrap(); assert_eq!(sub.id, id); } diff --git a/server/src/render/reddit.rs b/server/src/render/reddit.rs index 7f840aa33b734a31d8cf3341a0581c8bcb9bbcf3..595e202436040b0bfc419e68f083f94757ba5d0c 100644 --- a/server/src/render/reddit.rs +++ b/server/src/render/reddit.rs @@ -9,7 +9,7 @@ use crate::{ }; pub fn is_reddit_post(id: &ItemId) -> bool { - id.as_str().starts_with("reddit.com/") && id.as_str().contains("/comments/") + id.as_str().contains("reddit.com/") && id.as_str().contains("/comments/") } /// Post detail card (inside [`crate::fetch::html::entity_panel`]). diff --git a/server/src/state.rs b/server/src/state.rs index 78126f08d90f8069a586279d79258c27a9f9f7a4..513b329ef45fc332e63b3f8ed8498a1f59feb07c 100644 --- a/server/src/state.rs +++ b/server/src/state.rs @@ -217,15 +217,21 @@ impl AppState { ratio_right: i32, ) -> Result<(), String> { let ts = crate::html::now_ms(); - let vote = VoteData::from_recorded(ts, a, b, ratio_left, ratio_right) - .ok_or_else(|| "invalid vote: need two distinct non-empty items".to_string())?; + let a_raw = a.trim(); + let b_raw = b.trim(); + if a_raw.is_empty() || b_raw.is_empty() || a_raw == b_raw { + return Err("invalid vote: need two distinct non-empty items".to_string()); + } + // Validate items canonicalize (or are opaque keys) before append. + let _ = VoteData::from_recorded(ts, a_raw, b_raw, ratio_left, ratio_right) + .ok_or_else(|| "invalid vote: need two distinct parseable items".to_string())?; let event = Event::VoteRecorded { ts, - a: vote.a.as_str().to_string(), - b: vote.b.as_str().to_string(), - ratio_left: vote.ratio_left, - ratio_right: vote.ratio_right, + a: a_raw.to_string(), + b: b_raw.to_string(), + ratio_left, + ratio_right, scope: parent.as_str().to_string(), }; @@ -253,7 +259,7 @@ mod tests { let log = EventLog::new(log_path.to_string_lossy().into_owned()); let payload = json!({"kind":"t5","data":{"title":"Rust","display_name":"rust"}}); let event = Event::EntityImported { - id: "reddit.com/r/rust".into(), + id: "https://reddit.com/r/rust".into(), ts: 1, payload: payload.clone(), }; @@ -266,14 +272,14 @@ mod tests { .await .unwrap(); let tree = projection_store - .scope_tree(&ItemId::parse("reddit.com/r/rust").unwrap()) + .scope_tree(&ItemId::parse("https://reddit.com/r/rust").unwrap()) .unwrap(); let node = tree - .get(&ItemId::parse("reddit.com/r/rust").unwrap()) + .get(&ItemId::parse("https://reddit.com/r/rust").unwrap()) .unwrap(); assert_eq!(node.data.as_ref().unwrap().title, "Rust"); let stored = entity_store - .get(&ItemId::parse("reddit.com/r/rust").unwrap()) + .get(&ItemId::parse("https://reddit.com/r/rust").unwrap()) .unwrap() .unwrap(); assert_eq!(stored["data"]["display_name"], "rust"); @@ -289,13 +295,13 @@ mod tests { event_record( 1, Event::NodeEnsured { - id: "reddit.com/r/rust".into(), + id: "https://reddit.com/r/rust".into(), }, ), event_record( 2, Event::EntityImported { - id: "reddit.com/r/rust".into(), + id: "https://reddit.com/r/rust".into(), ts: 2, payload: payload.clone(), }, @@ -325,7 +331,7 @@ mod tests { &[event_record( 1, Event::NodeEnsured { - id: "reddit.com/r/stale".into(), + id: "https://reddit.com/r/stale".into(), }, )], ) @@ -352,11 +358,11 @@ mod tests { let root = tree.get(&ItemId::root()).unwrap(); assert!(root.children.contains(&ItemId::parse("alpha").unwrap())); assert!(projection_store - .load_node(&ItemId::parse("reddit.com/r/stale").unwrap()) + .load_node(&ItemId::parse("https://reddit.com/r/stale").unwrap()) .unwrap() .is_none()); let stored = entity_store - .get(&ItemId::parse("reddit.com/r/rust").unwrap()) + .get(&ItemId::parse("https://reddit.com/r/rust").unwrap()) .unwrap() .unwrap(); assert_eq!(stored["data"]["display_name"], "rust"); @@ -370,7 +376,7 @@ mod tests { log.append(&event_record( 1, Event::NodeEnsured { - id: "reddit.com/r/rust".into(), + id: "https://reddit.com/r/rust".into(), }, )) .await @@ -387,7 +393,7 @@ mod tests { &[event_record( 2, Event::NodeEnsured { - id: "reddit.com/r/rust".into(), + id: "https://reddit.com/r/rust".into(), }, )], ) @@ -454,7 +460,7 @@ mod tests { port: 0, }) .await; - let id = ItemId::parse("reddit.com/r/rust").unwrap(); + let id = ItemId::parse("https://reddit.com/r/rust").unwrap(); state.ensure_node(&id).await.unwrap(); @@ -465,11 +471,11 @@ mod tests { let projected = state.projection_store.load_tree().unwrap(); assert!(projected.get(&id).is_some()); let reddit = projected - .get(&ItemId::parse("reddit.com").unwrap()) + .get(&ItemId::from_url("https://reddit.com").unwrap()) .unwrap(); assert!(reddit .children - .contains(&ItemId::parse("reddit.com/r").unwrap())); + .contains(&ItemId::from_url("https://reddit.com/r").unwrap())); } #[tokio::test] @@ -510,7 +516,7 @@ mod tests { log.append(&event_record( 1, Event::NodeEnsured { - id: "reddit.com/r/rust".into(), + id: "https://reddit.com/r/rust".into(), }, )) .await @@ -518,7 +524,7 @@ mod tests { log.append(&event_record( 2, Event::NodeEnsured { - id: "reddit.com/r/python".into(), + id: "https://reddit.com/r/python".into(), }, )) .await @@ -544,13 +550,13 @@ mod tests { 2 ); let tree = second - .scope_tree(&ItemId::parse("reddit.com/r/rust").unwrap()) + .scope_tree(&ItemId::parse("https://reddit.com/r/rust").unwrap()) .unwrap(); assert!(tree - .get(&ItemId::parse("reddit.com/r/rust").unwrap()) + .get(&ItemId::parse("https://reddit.com/r/rust").unwrap()) .is_some()); assert!(tree - .get(&ItemId::parse("reddit.com/r/python").unwrap()) + .get(&ItemId::parse("https://reddit.com/r/python").unwrap()) .is_none()); } @@ -611,7 +617,7 @@ mod tests { #[test] fn parse_item_param_from_url() { let id = parse_item_param("https://reddit.com/r/rust"); - assert_eq!(id.as_str(), "reddit.com/r/rust"); + assert_eq!(id.as_str(), "https://reddit.com/r/rust"); } #[test] diff --git a/server/src/url_rules/engine.rs b/server/src/url_rules/engine.rs new file mode 100644 index 0000000000000000000000000000000000000000..e29b6b48c08deb7bffe031b1e542b1e25a7bef15 --- /dev/null +++ b/server/src/url_rules/engine.rs @@ -0,0 +1,187 @@ +//! Composable URL normalization primitives. + +use std::collections::HashMap; + +use url::Url; + +/// Mutable URL view used by rule combinators before serializing to a canonical string. +#[derive(Debug, Clone)] +pub struct ParsedUrl { + pub scheme: String, + pub host: String, + pub path_segments: Vec, + pub query: HashMap, + pub fragment: Option, +} + +impl ParsedUrl { + pub fn parse(raw: &str) -> Option { + let trimmed = raw.trim(); + if trimmed.is_empty() { + return None; + } + + let with_scheme = if trimmed.contains("://") { + trimmed.to_string() + } else if trimmed.starts_with("r/") || trimmed.starts_with("/r/") { + let rest = trimmed.trim_start_matches('/').trim_start_matches("r/"); + format!("https://reddit.com/r/{rest}") + } else if trimmed.contains('.') && !trimmed.starts_with('/') { + format!("https://{trimmed}") + } else { + trimmed.to_string() + }; + + let url = Url::parse(&with_scheme).ok()?; + let host = url.host_str()?.to_string(); + let path_segments: Vec = url + .path_segments() + .map(|segs| segs.filter(|s| !s.is_empty()).map(str::to_string).collect()) + .unwrap_or_default(); + + let mut query = HashMap::new(); + for (k, v) in url.query_pairs() { + query.insert(k.into_owned(), v.into_owned()); + } + + Some(Self { + scheme: url.scheme().to_string(), + path_segments, + query, + fragment: url.fragment().map(str::to_string), + host, + }) + } + + pub fn with_path_segments(&self, segments: &[String]) -> Self { + let mut u = self.clone(); + u.path_segments = segments.to_vec(); + u + } + + pub fn to_url(&self) -> Option { + let mut url = if self.path_segments.is_empty() { + Url::parse(&format!("{}://{}", self.scheme, self.host)).ok()? + } else { + let path = format!("/{}", self.path_segments.join("/")); + Url::parse(&format!("{}://{}{}", self.scheme, self.host, path)).ok()? + }; + if !self.query.is_empty() { + let mut pairs: Vec<_> = self.query.iter().collect(); + pairs.sort_by(|a, b| a.0.cmp(b.0)); + url.query_pairs_mut().clear(); + for (k, v) in pairs { + url.query_pairs_mut().append_pair(k, v); + } + } + if let Some(ref frag) = self.fragment { + url.set_fragment(Some(frag)); + } + Some(url) + } + + pub fn canonical_string(&self) -> Option { + let url = self.to_url()?; + let mut s = url.to_string(); + if self.path_segments.is_empty() { + s = s.trim_end_matches('/').to_string(); + } + Some(s) + } +} + +pub fn force_https(u: &mut ParsedUrl) { + if u.scheme == "http" { + u.scheme = "https".to_string(); + } +} + +pub fn drop_fragment(u: &mut ParsedUrl) { + u.fragment = None; +} + +pub fn strip_www(u: &mut ParsedUrl) { + if u.host.starts_with("www.") { + u.host = u.host[4..].to_string(); + } +} + +pub fn lowercase_host(u: &mut ParsedUrl) { + u.host = u.host.to_ascii_lowercase(); +} + +pub fn lowercase_path(u: &mut ParsedUrl) { + for seg in &mut u.path_segments { + *seg = seg.to_ascii_lowercase(); + } +} + +pub fn clear_query(u: &mut ParsedUrl) { + u.query.clear(); +} + +pub fn keep_only_query(u: &mut ParsedUrl, keys: &[&str]) { + u.query + .retain(|k, _| keys.iter().any(|want| want == &k.as_str())); +} + +pub fn strip_tracking_params(u: &mut ParsedUrl) { + u.query.retain(|k, _| { + let lower = k.to_ascii_lowercase(); + !(lower.starts_with("utm_") + || matches!( + lower.as_str(), + "fbclid" | "gclid" | "ref" | "ref_src" | "ref_source" | "mc_cid" | "mc_eid" + )) + }); +} + +pub fn truncate_after_segment(u: &mut ParsedUrl, name: &str, keep: usize) { + if let Some(i) = u.path_segments.iter().position(|s| s == name) { + let end = (i + 1 + keep).min(u.path_segments.len()); + u.path_segments.truncate(end); + } +} + +pub fn drop_listing_suffix(u: &mut ParsedUrl, suffixes: &[&str]) { + if u.path_segments.len() >= 3 && u.path_segments.first().map(String::as_str) == Some("r") { + if let Some(last) = u.path_segments.last() { + if suffixes.iter().any(|s| *s == last.as_str()) { + u.path_segments.pop(); + } + } + } +} + +pub fn normalize_reddit_host(u: &mut ParsedUrl) { + if matches!( + u.host.as_str(), + "old.reddit.com" | "new.reddit.com" | "www.reddit.com" + ) { + u.host = "reddit.com".to_string(); + } +} + +pub fn rewrite_youtu_be(u: &mut ParsedUrl) { + if u.host == "youtu.be" && u.path_segments.len() == 1 { + let id = u.path_segments[0].clone(); + u.host = "youtube.com".to_string(); + u.path_segments = vec!["watch".to_string()]; + u.query.insert("v".to_string(), id); + } +} + +pub fn rewrite_youtube_shorts(u: &mut ParsedUrl) { + if u.host == "youtube.com" && u.path_segments.first().map(String::as_str) == Some("shorts") { + if let Some(id) = u.path_segments.get(1).cloned() { + u.path_segments = vec!["watch".to_string()]; + u.query.insert("v".to_string(), id); + } + } +} + +pub fn normalize_youtube_host(u: &mut ParsedUrl) { + if matches!(u.host.as_str(), "m.youtube.com" | "www.youtube.com") { + u.host = "youtube.com".to_string(); + } +} diff --git a/server/src/url_rules/mod.rs b/server/src/url_rules/mod.rs new file mode 100644 index 0000000000000000000000000000000000000000..03d53bd3e82d704a01ba3fd8dd02b7d31422c0de --- /dev/null +++ b/server/src/url_rules/mod.rs @@ -0,0 +1,13 @@ +//! URL canonicalization and hierarchy rules for [`crate::path_types::ItemId`]. + +mod engine; +mod registry; + +pub use registry::{ + canonicalize_raw, looks_like_url, navigable_breadcrumbs, parent_url, resolve_id, CanonicalResult, +}; + +/// Resolve raw input to canonical URL. +pub fn resolve_canonical(raw: &str) -> Option { + canonicalize_raw(raw.trim()).map(|r| r.canonical) +} diff --git a/server/src/url_rules/registry.rs b/server/src/url_rules/registry.rs new file mode 100644 index 0000000000000000000000000000000000000000..14514e9af8385fb2b9b2f35eb9ee14d453d4b97c --- /dev/null +++ b/server/src/url_rules/registry.rs @@ -0,0 +1,235 @@ +//! Per-domain canonicalization and hierarchy rules. + +use std::collections::HashSet; + +use super::engine::{ + clear_query, drop_fragment, drop_listing_suffix, force_https, keep_only_query, lowercase_host, + lowercase_path, normalize_reddit_host, normalize_youtube_host, rewrite_youtu_be, + rewrite_youtube_shorts, strip_tracking_params, strip_www, truncate_after_segment, ParsedUrl, +}; + +/// Result of canonicalizing a raw URL string. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct CanonicalResult { + pub canonical: String, + /// When the input normalizes to a different string, the original is an alias. + pub alias_of: Option, +} + +fn apply_global(u: &mut ParsedUrl) { + force_https(u); + drop_fragment(u); + strip_www(u); + lowercase_host(u); + strip_tracking_params(u); +} + +fn normalize_reddit(u: &mut ParsedUrl) { + normalize_reddit_host(u); + lowercase_path(u); + truncate_after_segment(u, "comments", 1); + drop_listing_suffix(u, &["hot", "top", "new", "rising", "controversial"]); + clear_query(u); +} + +fn normalize_youtube(u: &mut ParsedUrl) { + rewrite_youtu_be(u); + normalize_youtube_host(u); + rewrite_youtube_shorts(u); + keep_only_query(u, &["v", "list"]); +} + +fn normalize_default(_u: &mut ParsedUrl) { + // Global rules only. +} + +fn domain_key(host: &str) -> &'static str { + if host == "reddit.com" || host.ends_with(".reddit.com") { + "reddit.com" + } else if host == "youtube.com" || host == "youtu.be" { + "youtube.com" + } else { + "default" + } +} + +fn normalize_for_host(u: &mut ParsedUrl) { + apply_global(u); + match domain_key(&u.host) { + "reddit.com" => normalize_reddit(u), + "youtube.com" => normalize_youtube(u), + _ => normalize_default(u), + } +} + +/// Structural path segments that must not become standalone tree nodes when more path follows. +fn structural_trailing(host: &str) -> &'static [&'static str] { + match domain_key(host) { + "reddit.com" => &["comments"], + _ => &[], + } +} + +/// Canonicalize a raw URL. Returns `None` if the input is not URL-like. +pub fn canonicalize_raw(raw: &str) -> Option { + let trimmed = raw.trim(); + if trimmed.is_empty() { + return None; + } + let mut u = ParsedUrl::parse(trimmed)?; + let input_snapshot = u.canonical_string()?; + normalize_for_host(&mut u); + let canonical = u.canonical_string()?; + let alias_of = if input_snapshot != canonical { + Some(trimmed.to_string()) + } else { + None + }; + Some(CanonicalResult { + canonical, + alias_of, + }) +} + +/// Resolve a stored or event id string to its canonical URL identity. +pub fn resolve_id(raw: &str) -> Option { + canonicalize_raw(raw).map(|r| r.canonical) +} + +/// Navigable ancestor URLs from domain root up to and including `canonical` (full URLs). +pub fn navigable_breadcrumbs(canonical: &str) -> Vec { + let Some(u) = ParsedUrl::parse(canonical) else { + return vec![canonical.to_string()]; + }; + let structural: HashSet<&str> = structural_trailing(&u.host).iter().copied().collect(); + let n = u.path_segments.len(); + let mut out = Vec::new(); + + // Domain root (no path segments). + if let Some(base) = u.with_path_segments(&[]).canonical_string() { + out.push(base); + } + + for i in 0..n { + let segs: Vec = u.path_segments[..=i].to_vec(); + let is_last = i == n - 1; + let seg = u.path_segments[i].as_str(); + if structural.contains(seg) && !is_last { + continue; + } + if let Some(url) = u.with_path_segments(&segs).canonical_string() { + if out.last() != Some(&url) { + out.push(url); + } + } + } + out +} + +/// Immediate parent scope URL, or `None` for tree root / opaque single-segment ids. +pub fn parent_url(canonical: &str) -> Option { + let crumbs = navigable_breadcrumbs(canonical); + if crumbs.len() <= 1 { + None + } else { + crumbs.get(crumbs.len() - 2).cloned() + } +} + +/// True when `raw` looks like a URL (has scheme or host-like shape). +pub fn looks_like_url(raw: &str) -> bool { + let t = raw.trim(); + t.contains("://") + || t.starts_with("r/") + || t.starts_with("/r/") + || (t.contains('.') && t.contains('/')) + || t.starts_with("reddit.com") + || t.starts_with("www.") + || t.starts_with("youtu.be/") +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn reddit_post_drops_slug_and_normalizes_host() { + let r = canonicalize_raw( + "https://old.reddit.com/r/AmItheAsshole/comments/1trnvdl/aita_for_cancelling/", + ) + .unwrap(); + assert_eq!( + r.canonical, + "https://reddit.com/r/amitheasshole/comments/1trnvdl" + ); + } + + #[test] + fn reddit_strips_query_and_listing() { + assert_eq!( + canonicalize_raw("https://www.reddit.com/r/rust/?sort=top") + .unwrap() + .canonical, + "https://reddit.com/r/rust" + ); + assert_eq!( + canonicalize_raw("https://www.reddit.com/r/programming/hot") + .unwrap() + .canonical, + "https://reddit.com/r/programming" + ); + } + + #[test] + fn reddit_short_path() { + assert_eq!( + canonicalize_raw("r/rust").unwrap().canonical, + "https://reddit.com/r/rust" + ); + } + + #[test] + fn reddit_breadcrumbs_skip_phantom_comments() { + let post = "https://reddit.com/r/aww/comments/1trnvdl"; + let crumbs = navigable_breadcrumbs(post); + assert!(!crumbs.iter().any(|c| c.ends_with("/comments"))); + assert_eq!( + crumbs.last().map(String::as_str), + Some(post) + ); + assert!(crumbs.contains(&"https://reddit.com/r/aww".to_string())); + } + + #[test] + fn reddit_parent_of_post_is_subreddit() { + assert_eq!( + parent_url("https://reddit.com/r/aww/comments/1trnvdl").as_deref(), + Some("https://reddit.com/r/aww") + ); + } + + #[test] + fn youtube_youtu_be_and_watch_same_canonical() { + let a = canonicalize_raw("https://youtu.be/dQw4w9WgXcQ").unwrap().canonical; + let b = canonicalize_raw("https://www.youtube.com/watch?v=dQw4w9WgXcQ&t=10").unwrap(); + assert_eq!(a, b.canonical); + assert_eq!(a, "https://youtube.com/watch?v=dQw4w9WgXcQ"); + } + + #[test] + fn legacy_schemeless_upgrades() { + assert_eq!( + canonicalize_raw("reddit.com/r/rust/comments/aaa/announcing_rust_199") + .unwrap() + .canonical, + "https://reddit.com/r/rust/comments/aaa" + ); + } + + #[test] + fn alias_recorded_when_input_differs() { + let r = canonicalize_raw("https://youtu.be/abc123").unwrap(); + assert_eq!(r.canonical, "https://youtube.com/watch?v=abc123"); + assert!(r.alias_of.is_some()); + } +}