diff options
| author | srdusr <[email protected]> | 2025-12-07 20:38:00 +0200 |
|---|---|---|
| committer | srdusr <[email protected]> | 2025-12-07 20:38:00 +0200 |
| commit | 5b1ea38522dbf6bf60db5a2270463de0c12d9de3 (patch) | |
| tree | c96e037c34230a05f2a457c5040843602dc7cf1f /crates/server/src | |
| parent | 24f1eb6cc611f458c45f0ac4046efce51211d7d3 (diff) | |
| download | typerpunk-5b1ea38522dbf6bf60db5a2270463de0c12d9de3.tar.gz typerpunk-5b1ea38522dbf6bf60db5a2270463de0c12d9de3.zip | |
Move the server to PostgreSQL, harden the lyrics proxy, add a hacking mode
PostgreSQL
- sqlx switched from the sqlite feature to postgres; the server now runs on
Postgres 18 and the SQLite file is gone.
- 95 placeholders renumbered from ? to $N.
- REAL widened to DOUBLE PRECISION: Postgres REAL is float4 and will not
decode into the f64 the code reads.
- flagged and is_bot are real BOOLEANs rather than 0/1 integers, with the
decode side reading bool.
- The leaderboard's derived table gained the alias Postgres requires, its
flag comparisons became boolean predicates, and INSERT OR IGNORE became
ON CONFLICT DO NOTHING.
- u32 binds cast to i64; Postgres has no unsigned integer types.
- Integration tests run against a real database - Postgres has no in-memory
mode - each in a throwaway schema, with search_path set per connection
because it is session state and the pool opens more than one.
- Timestamps stay TEXT for now and LISTEN/NOTIFY is still unused; both are
recorded in TODO-postgres.md rather than left implied.
Custom text and lyrics, checked rather than assumed
- Custom files never reach the server: they are read in the browser through
the File API, so there is no upload, no path handling and no remote file
inclusion to have. Verified by driving a hostile file - markup in the body
and in the filename - all the way onto the typing screen: it renders as
literal characters, no nodes are created, nothing executes, and the
filename is escaped in the attribution too.
- That test found a real regression: picking Custom from the new mode picker
selected it without ever starting it, so the mode was unstartable.
- /api/lyrics fixes its upstream host, so it cannot be pointed elsewhere, but
it was an unbounded relay: now rate limited per IP, with length caps on
artist and track and a ceiling on the response body it will read.
Hacking mode
- 22 single-line drills across recon, web, memory safety, exploit
development, crypto, post-exploitation and defence, each syntax
highlighted and each explaining what the line actually does.
All 19 modes verified to start, render and be typable.
Diffstat (limited to 'crates/server/src')
| -rw-r--r-- | crates/server/src/auth.rs | 30 | ||||
| -rw-r--r-- | crates/server/src/bot_results.rs | 8 | ||||
| -rw-r--r-- | crates/server/src/cosmetics.rs | 16 | ||||
| -rw-r--r-- | crates/server/src/friends.rs | 26 | ||||
| -rw-r--r-- | crates/server/src/lyrics.rs | 31 | ||||
| -rw-r--r-- | crates/server/src/main.rs | 55 | ||||
| -rw-r--r-- | crates/server/src/spotify.rs | 6 | ||||
| -rw-r--r-- | crates/server/src/state.rs | 11 | ||||
| -rw-r--r-- | crates/server/src/stats.rs | 20 |
9 files changed, 139 insertions, 64 deletions
diff --git a/crates/server/src/auth.rs b/crates/server/src/auth.rs index c1402d5..382dbfa 100644 --- a/crates/server/src/auth.rs +++ b/crates/server/src/auth.rs @@ -79,10 +79,10 @@ pub(crate) fn format_timestamp(dt: OffsetDateTime) -> String { dt.format(&time::format_description::well_known::Rfc3339).expect("valid RFC3339 timestamp") } -async fn create_session(db: &sqlx::SqlitePool, user_id: &str) -> Result<String, AppError> { +async fn create_session(db: &sqlx::PgPool, user_id: &str) -> Result<String, AppError> { let session_id = uuid::Uuid::new_v4().to_string(); let expires_at = OffsetDateTime::now_utc() + TimeDuration::days(SESSION_LIFETIME_DAYS); - sqlx::query("INSERT INTO sessions (id, user_id, expires_at) VALUES (?, ?, ?)") + sqlx::query("INSERT INTO sessions (id, user_id, expires_at) VALUES ($1, $2, $3)") .bind(&session_id) .bind(user_id) .bind(format_timestamp(expires_at)) @@ -102,12 +102,12 @@ fn session_cookie(id: String, secure: bool) -> Cookie<'static> { } /// Resolves the signed-in user from the session cookie, if any and unexpired. -pub async fn current_user(db: &sqlx::SqlitePool, jar: &CookieJar) -> Option<UserView> { +pub async fn current_user(db: &sqlx::PgPool, jar: &CookieJar) -> Option<UserView> { let session_id = jar.get(SESSION_COOKIE)?.value().to_string(); let row = sqlx::query( "SELECT users.id as id, users.username as username, sessions.expires_at as expires_at FROM sessions JOIN users ON users.id = sessions.user_id - WHERE sessions.id = ?", + WHERE sessions.id = $1", ) .bind(&session_id) .fetch_optional(db) @@ -139,7 +139,7 @@ async fn register( validate_username(&creds.username)?; validate_password(&creds.password)?; - let existing = sqlx::query("SELECT id FROM users WHERE username = ?") + let existing = sqlx::query("SELECT id FROM users WHERE username = $1") .bind(&creds.username) .fetch_optional(&state.db) .await?; @@ -149,7 +149,7 @@ async fn register( let user_id = uuid::Uuid::new_v4().to_string(); let password_hash = hash_password(&creds.password)?; - sqlx::query("INSERT INTO users (id, username, password_hash, created_at) VALUES (?, ?, ?, ?)") + sqlx::query("INSERT INTO users (id, username, password_hash, created_at) VALUES ($1, $2, $3, $4)") .bind(&user_id) .bind(&creds.username) .bind(&password_hash) @@ -166,8 +166,8 @@ async fn register( /// cookie-session login and the CLI-style token login below, so a /// nonexistent username always takes the same dummy-hash timing path /// regardless of which endpoint is asking. -async fn authenticate(db: &sqlx::SqlitePool, creds: &Credentials) -> Result<UserView, AppError> { - let row = sqlx::query("SELECT id, username, password_hash FROM users WHERE username = ?") +async fn authenticate(db: &sqlx::PgPool, creds: &Credentials) -> Result<UserView, AppError> { + let row = sqlx::query("SELECT id, username, password_hash FROM users WHERE username = $1") .bind(&creds.username) .fetch_optional(db) .await?; @@ -223,7 +223,7 @@ async fn issue_token( } let user = authenticate(&state.db, &creds).await?; let token = uuid::Uuid::new_v4().to_string(); - sqlx::query("INSERT INTO api_tokens (token, user_id, created_at) VALUES (?, ?, ?)") + sqlx::query("INSERT INTO api_tokens (token, user_id, created_at) VALUES ($1, $2, $3)") .bind(&token) .bind(&user.id) .bind(format_timestamp(OffsetDateTime::now_utc())) @@ -234,11 +234,11 @@ async fn issue_token( /// Resolves a signed-in user from a bearer token (see api_tokens above), /// for clients with no cookie jar. -pub async fn user_from_token(db: &sqlx::SqlitePool, token: &str) -> Option<UserView> { +pub async fn user_from_token(db: &sqlx::PgPool, token: &str) -> Option<UserView> { let row = sqlx::query( "SELECT users.id as id, users.username as username FROM api_tokens JOIN users ON users.id = api_tokens.user_id - WHERE api_tokens.token = ?", + WHERE api_tokens.token = $1", ) .bind(token) .fetch_optional(db) @@ -262,11 +262,11 @@ const PRESENCE_WRITE_INTERVAL_SECS: i64 = 60; /// authenticated request: the WHERE clause skips the write unless the stored /// value is already stale, so it is one no-op UPDATE a minute per active user /// rather than one per request. -async fn touch_last_seen(db: &sqlx::SqlitePool, user_id: &str) { +async fn touch_last_seen(db: &sqlx::PgPool, user_id: &str) { let now = OffsetDateTime::now_utc(); let cutoff = now - TimeDuration::seconds(PRESENCE_WRITE_INTERVAL_SECS); let _ = sqlx::query( - "UPDATE users SET last_seen = ? WHERE id = ? AND (last_seen IS NULL OR last_seen < ?)", + "UPDATE users SET last_seen = $1 WHERE id = $2 AND (last_seen IS NULL OR last_seen < $3)", ) .bind(format_timestamp(now)) .bind(user_id) @@ -276,7 +276,7 @@ async fn touch_last_seen(db: &sqlx::SqlitePool, user_id: &str) { } pub async fn current_user_or_token( - db: &sqlx::SqlitePool, + db: &sqlx::PgPool, jar: &CookieJar, headers: &axum::http::HeaderMap, ) -> Option<UserView> { @@ -296,7 +296,7 @@ pub async fn current_user_or_token( async fn logout(State(state): State<Arc<AppState>>, jar: CookieJar) -> Result<impl IntoResponse, AppError> { if let Some(cookie) = jar.get(SESSION_COOKIE) { - sqlx::query("DELETE FROM sessions WHERE id = ?") + sqlx::query("DELETE FROM sessions WHERE id = $1") .bind(cookie.value()) .execute(&state.db) .await?; diff --git a/crates/server/src/bot_results.rs b/crates/server/src/bot_results.rs index 255270e..c777d9c 100644 --- a/crates/server/src/bot_results.rs +++ b/crates/server/src/bot_results.rs @@ -68,7 +68,7 @@ async fn ensure_bot_users(state: &AppState) -> Result<(), sqlx::Error> { // into, and the auth path compares against a hash that cannot match. sqlx::query( "INSERT INTO users (id, username, password_hash, created_at, is_bot) - VALUES (?, ?, '!', ?, 1) + VALUES ($1, $2, '!', $3, TRUE) ON CONFLICT(username) DO NOTHING", ) .bind(format!("bot-{name}")) @@ -95,7 +95,7 @@ async fn post_one_result(state: &AppState) -> Result<(), sqlx::Error> { sqlx::query( "INSERT INTO test_results (id, user_id, mode_key, wpm, raw_wpm, accuracy, time_seconds, created_at, device_type, flagged) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, 'desktop', 0)", + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, 'desktop', FALSE)", ) .bind(uuid::Uuid::new_v4().to_string()) .bind(format!("bot-{name}")) @@ -122,7 +122,7 @@ async fn seed_backlog(state: &AppState) -> Result<(), sqlx::Error> { // leave every other mode permanently empty. This also means a mode // added later gets seeded on the next start. let existing: i64 = sqlx::query_scalar( - "SELECT COUNT(*) FROM test_results WHERE user_id LIKE 'bot-%' AND mode_key = ?", + "SELECT COUNT(*) FROM test_results WHERE user_id LIKE 'bot-%' AND mode_key = $1", ) .bind(*mode) .fetch_one(&state.db) @@ -144,7 +144,7 @@ async fn seed_backlog(state: &AppState) -> Result<(), sqlx::Error> { let created = OffsetDateTime::now_utc() - time::Duration::hours(age_hours); sqlx::query( "INSERT INTO test_results (id, user_id, mode_key, wpm, raw_wpm, accuracy, time_seconds, created_at, device_type, flagged) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, 'desktop', 0)", + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, 'desktop', FALSE)", ) .bind(uuid::Uuid::new_v4().to_string()) .bind(format!("bot-{name}")) diff --git a/crates/server/src/cosmetics.rs b/crates/server/src/cosmetics.rs index b3eee28..f128095 100644 --- a/crates/server/src/cosmetics.rs +++ b/crates/server/src/cosmetics.rs @@ -56,13 +56,13 @@ struct MyCosmetics { async fn my_cosmetics(State(state): State<Arc<AppState>>, jar: CookieJar) -> Result<impl IntoResponse, AppError> { let user = current_user(&state.db, &jar).await.ok_or(AppError::Unauthorized)?; - let owned_rows = sqlx::query("SELECT cosmetic_id FROM user_cosmetics WHERE user_id = ?") + let owned_rows = sqlx::query("SELECT cosmetic_id FROM user_cosmetics WHERE user_id = $1") .bind(&user.id) .fetch_all(&state.db) .await?; let owned: Vec<String> = owned_rows.into_iter().filter_map(|r| r.try_get("cosmetic_id").ok()).collect(); - let equip_row = sqlx::query("SELECT equipped_caret, equipped_flair FROM users WHERE id = ?") + let equip_row = sqlx::query("SELECT equipped_caret, equipped_flair FROM users WHERE id = $1") .bind(&user.id) .fetch_one(&state.db) .await?; @@ -88,7 +88,7 @@ async fn my_cosmetics(State(state): State<Arc<AppState>>, jar: CookieJar) -> Res async fn purchase(State(state): State<Arc<AppState>>, jar: CookieJar, Path(cosmetic_id): Path<String>) -> Result<impl IntoResponse, AppError> { let user = current_user(&state.db, &jar).await.ok_or(AppError::Unauthorized)?; - let exists = sqlx::query("SELECT id FROM cosmetics WHERE id = ?") + let exists = sqlx::query("SELECT id FROM cosmetics WHERE id = $1") .bind(&cosmetic_id) .fetch_optional(&state.db) .await?; @@ -96,7 +96,9 @@ async fn purchase(State(state): State<Arc<AppState>>, jar: CookieJar, Path(cosme return Err(AppError::NotFound); } - sqlx::query("INSERT OR IGNORE INTO user_cosmetics (user_id, cosmetic_id, acquired_at) VALUES (?, ?, ?)") + // Postgres spells SQLite's INSERT OR IGNORE as an explicit conflict + // target; the pair is the table's primary key. + sqlx::query("INSERT INTO user_cosmetics (user_id, cosmetic_id, acquired_at) VALUES ($1, $2, $3) ON CONFLICT (user_id, cosmetic_id) DO NOTHING") .bind(&user.id) .bind(&cosmetic_id) .bind(format_timestamp(OffsetDateTime::now_utc())) @@ -112,7 +114,7 @@ async fn equip(State(state): State<Arc<AppState>>, jar: CookieJar, Path(cosmetic let row = sqlx::query( "SELECT cosmetics.category as category FROM cosmetics JOIN user_cosmetics ON user_cosmetics.cosmetic_id = cosmetics.id - WHERE cosmetics.id = ? AND user_cosmetics.user_id = ?", + WHERE cosmetics.id = $1 AND user_cosmetics.user_id = $2", ) .bind(&cosmetic_id) .bind(&user.id) @@ -129,7 +131,7 @@ async fn equip(State(state): State<Arc<AppState>>, jar: CookieJar, Path(cosmetic // Column name comes from a hardcoded match above, never from request // input, so this is safe despite not being a bind parameter. - let query = format!("UPDATE users SET {column} = ? WHERE id = ?"); + let query = format!("UPDATE users SET {column} = $1 WHERE id = $2"); sqlx::query(&query).bind(&cosmetic_id).bind(&user.id).execute(&state.db).await?; Ok(axum::http::StatusCode::NO_CONTENT) @@ -147,7 +149,7 @@ async fn unequip(State(state): State<Arc<AppState>>, jar: CookieJar, Json(body): "flair" => "equipped_flair", _ => return Err(AppError::InvalidInput("category must be 'caret' or 'flair'".into())), }; - let query = format!("UPDATE users SET {column} = NULL WHERE id = ?"); + let query = format!("UPDATE users SET {column} = NULL WHERE id = $1"); sqlx::query(&query).bind(&user.id).execute(&state.db).await?; Ok(axum::http::StatusCode::NO_CONTENT) } diff --git a/crates/server/src/friends.rs b/crates/server/src/friends.rs index 6621def..48cd549 100644 --- a/crates/server/src/friends.rs +++ b/crates/server/src/friends.rs @@ -37,13 +37,13 @@ struct FriendsList { outgoing_requests: Vec<FriendEntry>, } -fn row_to_entry(row: &sqlx::sqlite::SqliteRow, friendship_id_col: &str, user_id_col: &str, username_col: &str) -> FriendEntry { +fn row_to_entry(row: &sqlx::postgres::PgRow, friendship_id_col: &str, user_id_col: &str, username_col: &str) -> FriendEntry { FriendEntry { friendship_id: row.try_get(friendship_id_col).unwrap_or_default(), user_id: row.try_get(user_id_col).unwrap_or_default(), username: row.try_get(username_col).unwrap_or_default(), // Absent on the pending-request queries, which do not select it. - online: row.try_get::<i64, _>("online").unwrap_or(0) != 0, + online: row.try_get::<bool, _>("online").unwrap_or(false), } } @@ -60,10 +60,10 @@ async fn list_friends(State(state): State<Arc<AppState>>, jar: CookieJar, header ); let accepted = sqlx::query( "SELECT friendships.id as friendship_id, users.id as user_id, users.username as username, - (users.last_seen IS NOT NULL AND users.last_seen > ?) as online + (users.last_seen IS NOT NULL AND users.last_seen > $1) as online FROM friendships - JOIN users ON users.id = CASE WHEN friendships.requester_id = ? THEN friendships.addressee_id ELSE friendships.requester_id END - WHERE friendships.status = 'accepted' AND (friendships.requester_id = ? OR friendships.addressee_id = ?)", + JOIN users ON users.id = CASE WHEN friendships.requester_id = $2 THEN friendships.addressee_id ELSE friendships.requester_id END + WHERE friendships.status = 'accepted' AND (friendships.requester_id = $3 OR friendships.addressee_id = $4)", ) .bind(&presence_cutoff) .bind(&user.id).bind(&user.id).bind(&user.id) @@ -73,7 +73,7 @@ async fn list_friends(State(state): State<Arc<AppState>>, jar: CookieJar, header let incoming = sqlx::query( "SELECT friendships.id as friendship_id, users.id as user_id, users.username as username FROM friendships JOIN users ON users.id = friendships.requester_id - WHERE friendships.status = 'pending' AND friendships.addressee_id = ?", + WHERE friendships.status = 'pending' AND friendships.addressee_id = $1", ) .bind(&user.id) .fetch_all(&state.db) @@ -82,7 +82,7 @@ async fn list_friends(State(state): State<Arc<AppState>>, jar: CookieJar, header let outgoing = sqlx::query( "SELECT friendships.id as friendship_id, users.id as user_id, users.username as username FROM friendships JOIN users ON users.id = friendships.addressee_id - WHERE friendships.status = 'pending' AND friendships.requester_id = ?", + WHERE friendships.status = 'pending' AND friendships.requester_id = $1", ) .bind(&user.id) .fetch_all(&state.db) @@ -108,7 +108,7 @@ async fn send_request( ) -> Result<impl IntoResponse, AppError> { let user = current_user_or_token(&state.db, &jar, &headers).await.ok_or(AppError::Unauthorized)?; - let target = sqlx::query("SELECT id FROM users WHERE username = ?") + let target = sqlx::query("SELECT id FROM users WHERE username = $1") .bind(&body.username) .fetch_optional(&state.db) .await? @@ -121,7 +121,7 @@ async fn send_request( let existing = sqlx::query( "SELECT id, requester_id, status FROM friendships - WHERE (requester_id = ? AND addressee_id = ?) OR (requester_id = ? AND addressee_id = ?)", + WHERE (requester_id = $1 AND addressee_id = $2) OR (requester_id = $3 AND addressee_id = $4)", ) .bind(&user.id).bind(&target_id) .bind(&target_id).bind(&user.id) @@ -138,7 +138,7 @@ async fn send_request( // creating a second, redundant pending row in the opposite direction. if requester_id == target_id { let id: String = row.try_get("id").map_err(|e| AppError::Internal(e.into()))?; - sqlx::query("UPDATE friendships SET status = 'accepted' WHERE id = ?") + sqlx::query("UPDATE friendships SET status = 'accepted' WHERE id = $1") .bind(&id) .execute(&state.db) .await?; @@ -147,7 +147,7 @@ async fn send_request( return Err(AppError::InvalidInput("request already pending".into())); } - sqlx::query("INSERT INTO friendships (id, requester_id, addressee_id, status, created_at) VALUES (?, ?, ?, 'pending', ?)") + sqlx::query("INSERT INTO friendships (id, requester_id, addressee_id, status, created_at) VALUES ($1, $2, $3, 'pending', $4)") .bind(uuid::Uuid::new_v4().to_string()) .bind(&user.id) .bind(&target_id) @@ -167,7 +167,7 @@ async fn accept_request( let user = current_user_or_token(&state.db, &jar, &headers).await.ok_or(AppError::Unauthorized)?; let result = sqlx::query( - "UPDATE friendships SET status = 'accepted' WHERE id = ? AND addressee_id = ? AND status = 'pending'", + "UPDATE friendships SET status = 'accepted' WHERE id = $1 AND addressee_id = $2 AND status = 'pending'", ) .bind(&id) .bind(&user.id) @@ -191,7 +191,7 @@ async fn remove_friendship( ) -> Result<impl IntoResponse, AppError> { let user = current_user_or_token(&state.db, &jar, &headers).await.ok_or(AppError::Unauthorized)?; - let result = sqlx::query("DELETE FROM friendships WHERE id = ? AND (requester_id = ? OR addressee_id = ?)") + let result = sqlx::query("DELETE FROM friendships WHERE id = $1 AND (requester_id = $2 OR addressee_id = $3)") .bind(&id) .bind(&user.id) .bind(&user.id) diff --git a/crates/server/src/lyrics.rs b/crates/server/src/lyrics.rs index d7c8710..332d9eb 100644 --- a/crates/server/src/lyrics.rs +++ b/crates/server/src/lyrics.rs @@ -1,6 +1,6 @@ use crate::error::AppError; use crate::state::AppState; -use axum::extract::{Query, State}; +use axum::extract::{ConnectInfo, Query, State}; use axum::response::IntoResponse; use axum::routing::get; use axum::{Json, Router}; @@ -37,10 +37,32 @@ pub struct LyricsResult { pub plain: Option<String>, } -async fn get_lyrics(State(state): State<Arc<AppState>>, Query(q): Query<LyricsQuery>) -> Result<impl IntoResponse, AppError> { +/// Beyond this, a value is not a track or artist name - it is someone using +/// this endpoint to push a large query at lrclib. +const MAX_FIELD_LEN: usize = 200; + +/// The upstream body is read into memory, so it needs a ceiling that does not +/// depend on the third party behaving. Lyrics for a song are a few kilobytes. +const MAX_RESPONSE_BYTES: usize = 256 * 1024; + +async fn get_lyrics( + State(state): State<Arc<AppState>>, + ConnectInfo(addr): ConnectInfo<std::net::SocketAddr>, + Query(q): Query<LyricsQuery>, +) -> Result<impl IntoResponse, AppError> { + // This endpoint makes an outbound request on the caller's behalf. The + // destination is fixed, so it cannot be pointed at anything else, but + // without a limit it is still a free relay for hammering lrclib from our + // address rather than the caller's. + if !state.lyrics_rate_limiter.check(addr.ip()) { + return Err(AppError::RateLimited); + } if q.artist.trim().is_empty() || q.track.trim().is_empty() { return Err(AppError::InvalidInput("artist and track are required".into())); } + if q.artist.len() > MAX_FIELD_LEN || q.track.len() > MAX_FIELD_LEN { + return Err(AppError::InvalidInput("artist and track are too long".into())); + } let mut req = state .http @@ -51,6 +73,11 @@ async fn get_lyrics(State(state): State<Arc<AppState>>, Query(q): Query<LyricsQu } let res = req.send().await.map_err(|e| AppError::Internal(e.into()))?; + if let Some(len) = res.content_length() { + if len as usize > MAX_RESPONSE_BYTES { + return Err(AppError::Internal(anyhow::anyhow!("lrclib response too large"))); + } + } if res.status() == reqwest::StatusCode::NOT_FOUND { return Err(AppError::NotFound); } diff --git a/crates/server/src/main.rs b/crates/server/src/main.rs index 0d1c980..e9133b9 100644 --- a/crates/server/src/main.rs +++ b/crates/server/src/main.rs @@ -79,13 +79,18 @@ async fn main() -> anyhow::Result<()> { .with_env_filter(tracing_subscriber::EnvFilter::try_from_default_env().unwrap_or_else(|_| "info".into())) .init(); - let database_url = std::env::var("DATABASE_URL").unwrap_or_else(|_| "sqlite://typerpunk.db".to_string()); + let database_url = std::env::var("DATABASE_URL") + .unwrap_or_else(|_| "postgres://typerpunk:[email protected]/typerpunk".to_string()); let cookie_secure = std::env::var("COOKIE_SECURE").map(|v| v == "1" || v == "true").unwrap_or(false); let frontend_origin = std::env::var("FRONTEND_ORIGIN").unwrap_or_else(|_| "http://localhost:4173".to_string()); let port: u16 = std::env::var("PORT").ok().and_then(|p| p.parse().ok()).unwrap_or(8787); - let connect_options = sqlx::sqlite::SqliteConnectOptions::from_str(&database_url)?.create_if_missing(true); - let db = sqlx::sqlite::SqlitePoolOptions::new().max_connections(10).connect_with(connect_options).await?; + // No create_if_missing: Postgres databases are created by an administrator, + // not by the application on first connect the way a SQLite file was. + let db = sqlx::postgres::PgPoolOptions::new() + .max_connections(10) + .connect(&database_url) + .await?; sqlx::migrate!("./migrations").run(&db).await?; if !cookie_secure { @@ -134,11 +139,47 @@ mod tests { use state::SpotifyConfig; async fn spawn_test_server() -> String { - let db = sqlx::sqlite::SqlitePoolOptions::new() - .max_connections(1) - .connect("sqlite::memory:") + // Postgres has no in-memory mode, so these run against a real database. + // Each test gets its own schema inside it: a shared one would let two + // tests running concurrently see each other's rows, and DROP SCHEMA + // cleans up without coordinating table lists. + let url = std::env::var("TEST_DATABASE_URL").unwrap_or_else(|_| { + "postgres://typerpunk:[email protected]/typerpunk_test".to_string() + }); + let schema = format!("t{}", uuid::Uuid::new_v4().simple()); + + // Created over a one-shot connection first, because the pool below + // pins every connection to this schema and it has to exist by then. + { + let setup = sqlx::postgres::PgPoolOptions::new() + .max_connections(1) + .connect(&url) + .await + .expect("failed to connect to the test database - is Postgres running, and does typerpunk_test exist?"); + sqlx::query(&format!("CREATE SCHEMA {schema}")) + .execute(&setup) + .await + .expect("failed to create test schema"); + setup.close().await; + } + + // search_path is per-session, so it is set on every connection the + // pool opens rather than once on whichever one happened to be first. + let schema_for_hook = schema.clone(); + let db = sqlx::postgres::PgPoolOptions::new() + .max_connections(2) + .after_connect(move |conn, _meta| { + let schema = schema_for_hook.clone(); + Box::pin(async move { + sqlx::query(&format!("SET search_path TO {schema}")) + .execute(conn) + .await + .map(|_| ()) + }) + }) + .connect(&url) .await - .expect("failed to open in-memory test database"); + .expect("failed to connect to the test database"); sqlx::migrate!("./migrations").run(&db).await.expect("failed to run migrations"); let app_state = Arc::new(AppState::new( diff --git a/crates/server/src/spotify.rs b/crates/server/src/spotify.rs index 9076afd..0ed652e 100644 --- a/crates/server/src/spotify.rs +++ b/crates/server/src/spotify.rs @@ -110,7 +110,7 @@ async fn callback( let expires_at = format_timestamp(OffsetDateTime::now_utc() + TimeDuration::seconds(token.expires_in)); sqlx::query( - "INSERT INTO spotify_tokens (user_id, access_token, refresh_token, expires_at) VALUES (?, ?, ?, ?) + "INSERT INTO spotify_tokens (user_id, access_token, refresh_token, expires_at) VALUES ($1, $2, $3, $4) ON CONFLICT(user_id) DO UPDATE SET access_token = excluded.access_token, refresh_token = excluded.refresh_token, expires_at = excluded.expires_at", ) .bind(&user.id) @@ -125,7 +125,7 @@ async fn callback( } async fn get_valid_access_token(state: &AppState, user_id: &str) -> Result<Option<String>, AppError> { - let row = sqlx::query("SELECT access_token, refresh_token, expires_at FROM spotify_tokens WHERE user_id = ?") + let row = sqlx::query("SELECT access_token, refresh_token, expires_at FROM spotify_tokens WHERE user_id = $1") .bind(user_id) .fetch_optional(&state.db) .await?; @@ -159,7 +159,7 @@ async fn get_valid_access_token(state: &AppState, user_id: &str) -> Result<Optio let new_refresh_token = refreshed.refresh_token.unwrap_or(refresh_token); let new_expires_at = format_timestamp(OffsetDateTime::now_utc() + TimeDuration::seconds(refreshed.expires_in)); - sqlx::query("UPDATE spotify_tokens SET access_token = ?, refresh_token = ?, expires_at = ? WHERE user_id = ?") + sqlx::query("UPDATE spotify_tokens SET access_token = $1, refresh_token = $2, expires_at = $3 WHERE user_id = $4") .bind(&refreshed.access_token) .bind(&new_refresh_token) .bind(&new_expires_at) diff --git a/crates/server/src/state.rs b/crates/server/src/state.rs index 7244955..fac6f29 100644 --- a/crates/server/src/state.rs +++ b/crates/server/src/state.rs @@ -1,7 +1,7 @@ use crate::multiplayer::{new_registry, RoomRegistry}; use crate::rate_limit::RateLimiter; use reqwest::Client; -use sqlx::SqlitePool; +use sqlx::PgPool; use std::time::Duration; #[derive(Clone, Default)] @@ -28,13 +28,17 @@ pub struct RaceText { #[derive(Clone)] pub struct AppState { - pub db: SqlitePool, + pub db: PgPool, pub auth_rate_limiter: RateLimiter, /// Keyed by user id, not IP - this guards an authenticated endpoint /// (stats submission) against a single compromised/scripted account /// hammering it, which an IP-keyed limiter wouldn't catch behind NAT or /// a VPN and would over-punish for a shared IP. pub stats_rate_limiter: RateLimiter<String>, + /// /api/lyrics forwards to a third party on the caller's behalf, so it is + /// an open proxy unless it is bounded. Keyed by IP, since the endpoint is + /// reachable without an account. + pub lyrics_rate_limiter: RateLimiter, /// Set from COOKIE_SECURE. Off for plain-HTTP local dev, must be on /// behind TLS in production or browsers silently drop the cookie. pub cookie_secure: bool, @@ -50,7 +54,7 @@ pub struct AppState { impl AppState { pub fn new( - db: SqlitePool, + db: PgPool, cookie_secure: bool, race_texts: Vec<RaceText>, spotify: SpotifyConfig, @@ -63,6 +67,7 @@ impl AppState { // seconds; 60 submissions in 5 minutes is generous headroom for // rapid Words-10 sessions while still capping scripted spam. stats_rate_limiter: RateLimiter::new(60, Duration::from_secs(5 * 60)), + lyrics_rate_limiter: RateLimiter::new(30, Duration::from_secs(60)), cookie_secure, rooms: new_registry(), race_texts, diff --git a/crates/server/src/stats.rs b/crates/server/src/stats.rs index a7d3a70..931e096 100644 --- a/crates/server/src/stats.rs +++ b/crates/server/src/stats.rs @@ -75,7 +75,7 @@ async fn submit_result( sqlx::query( "INSERT INTO test_results (id, user_id, mode_key, wpm, raw_wpm, accuracy, time_seconds, created_at, device_type, flagged) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)", ) .bind(uuid::Uuid::new_v4().to_string()) .bind(&user.id) @@ -120,7 +120,7 @@ async fn fetch_stats_summary(state: &AppState, user_id: &str) -> Result<MyStats, COALESCE(AVG(wpm), 0) as average_wpm, COALESCE(AVG(accuracy), 0) as average_accuracy, COALESCE(MAX(wpm), 0) as best_wpm - FROM test_results WHERE user_id = ?", + FROM test_results WHERE user_id = $1", ) .bind(user_id) .fetch_one(&state.db) @@ -133,7 +133,7 @@ async fn fetch_stats_summary(state: &AppState, user_id: &str) -> Result<MyStats, "SELECT mode_key, wpm, created_at FROM ( SELECT mode_key, wpm, created_at, ROW_NUMBER() OVER (PARTITION BY mode_key ORDER BY wpm DESC) as rn - FROM test_results WHERE user_id = ? + FROM test_results WHERE user_id = $1 ) WHERE rn = 1", ) .bind(user_id) @@ -187,7 +187,7 @@ async fn public_profile(State(state): State<Arc<AppState>>, Path(username): Path let row = sqlx::query( "SELECT users.id as id, users.created_at as created_at, cosmetics.value as flair_value FROM users LEFT JOIN cosmetics ON cosmetics.id = users.equipped_flair - WHERE users.username = ?", + WHERE users.username = $1", ) .bind(&username) .fetch_optional(&state.db) @@ -252,15 +252,15 @@ async fn leaderboard(State(state): State<Arc<AppState>>, Query(q): Query<Leaderb FROM test_results JOIN users ON users.id = test_results.user_id LEFT JOIN cosmetics ON cosmetics.id = users.equipped_flair - WHERE test_results.mode_key = ? AND test_results.flagged = 0 - AND (? = 0 OR test_results.device_type = 'desktop') - ) WHERE rn = 1 + WHERE test_results.mode_key = $1 AND NOT test_results.flagged + AND (NOT $2 OR test_results.device_type = 'desktop') + ) AS ranked WHERE rn = 1 ORDER BY wpm DESC - LIMIT ?", + LIMIT $3", ) .bind(&q.mode) .bind(desktop_only) - .bind(limit) + .bind(limit as i64) .fetch_all(&state.db) .await?; @@ -273,7 +273,7 @@ async fn leaderboard(State(state): State<Arc<AppState>>, Query(q): Query<Leaderb date: row.try_get("created_at").unwrap_or_default(), device_type: row.try_get("device_type").unwrap_or_default(), flair: row.try_get::<Option<String>, _>("flair_value").unwrap_or(None), - is_bot: row.try_get::<i64, _>("is_bot").unwrap_or(0) != 0, + is_bot: row.try_get::<bool, _>("is_bot").unwrap_or(false), }) .collect(); |