srdusr
aboutsummaryrefslogtreecommitdiffstats
path: root/crates/server
diff options
context:
space:
mode:
Diffstat (limited to 'crates/server')
-rw-r--r--crates/server/.env.example9
-rw-r--r--crates/server/migrations/0002_test_results.sql8
-rw-r--r--crates/server/migrations/0005_fairness.sql2
-rw-r--r--crates/server/migrations/0009_bot_players.sql2
-rw-r--r--crates/server/src/auth.rs30
-rw-r--r--crates/server/src/bot_results.rs8
-rw-r--r--crates/server/src/cosmetics.rs16
-rw-r--r--crates/server/src/friends.rs26
-rw-r--r--crates/server/src/lyrics.rs31
-rw-r--r--crates/server/src/main.rs55
-rw-r--r--crates/server/src/spotify.rs6
-rw-r--r--crates/server/src/state.rs11
-rw-r--r--crates/server/src/stats.rs20
13 files changed, 153 insertions, 71 deletions
diff --git a/crates/server/.env.example b/crates/server/.env.example
index 6dabec5..8485826 100644
--- a/crates/server/.env.example
+++ b/crates/server/.env.example
@@ -1,4 +1,11 @@
-DATABASE_URL=sqlite://typerpunk.db
+# Postgres. Create the database and role first:
+# sudo -u postgres createuser --pwprompt typerpunk
+# sudo -u postgres createdb -O typerpunk typerpunk
+DATABASE_URL=postgres://typerpunk:[email protected]/typerpunk
+
+# The integration tests need their own database; each test creates a throwaway
+# schema inside it.
+TEST_DATABASE_URL=postgres://typerpunk:[email protected]/typerpunk_test
PORT=8787
FRONTEND_ORIGIN=http://localhost:4173
diff --git a/crates/server/migrations/0002_test_results.sql b/crates/server/migrations/0002_test_results.sql
index aa876b5..2a4cf69 100644
--- a/crates/server/migrations/0002_test_results.sql
+++ b/crates/server/migrations/0002_test_results.sql
@@ -2,10 +2,10 @@ CREATE TABLE test_results (
id TEXT PRIMARY KEY,
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
mode_key TEXT NOT NULL,
- wpm REAL NOT NULL,
- raw_wpm REAL NOT NULL,
- accuracy REAL NOT NULL,
- time_seconds REAL NOT NULL,
+ wpm DOUBLE PRECISION NOT NULL,
+ raw_wpm DOUBLE PRECISION NOT NULL,
+ accuracy DOUBLE PRECISION NOT NULL,
+ time_seconds DOUBLE PRECISION NOT NULL,
created_at TEXT NOT NULL
);
diff --git a/crates/server/migrations/0005_fairness.sql b/crates/server/migrations/0005_fairness.sql
index 28152d6..6d55b52 100644
--- a/crates/server/migrations/0005_fairness.sql
+++ b/crates/server/migrations/0005_fairness.sql
@@ -1,2 +1,2 @@
ALTER TABLE test_results ADD COLUMN device_type TEXT NOT NULL DEFAULT 'desktop';
-ALTER TABLE test_results ADD COLUMN flagged INTEGER NOT NULL DEFAULT 0;
+ALTER TABLE test_results ADD COLUMN flagged BOOLEAN NOT NULL DEFAULT FALSE;
diff --git a/crates/server/migrations/0009_bot_players.sql b/crates/server/migrations/0009_bot_players.sql
index 790eb85..72516c5 100644
--- a/crates/server/migrations/0009_bot_players.sql
+++ b/crates/server/migrations/0009_bot_players.sql
@@ -3,6 +3,6 @@
- the same cosmetics. The flag exists so the client can label them, because a
- board that mixes synthetic scores into human ones without saying so is
- telling its users something untrue.
-ALTER TABLE users ADD COLUMN is_bot INTEGER NOT NULL DEFAULT 0;
+ALTER TABLE users ADD COLUMN is_bot BOOLEAN NOT NULL DEFAULT FALSE;
CREATE INDEX idx_users_is_bot ON users(is_bot);
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();