diff --git a/spotify_proxy/src/spotify_proxy.gleam b/spotify_proxy/src/spotify_proxy.gleam index e5b8d12..2ee0749 100644 --- a/spotify_proxy/src/spotify_proxy.gleam +++ b/spotify_proxy/src/spotify_proxy.gleam @@ -6,6 +6,7 @@ import gleam/otp/supervision import gleam/result import logging import spotify_proxy/api +import spotify_proxy/ratelimiter import spotify_proxy/spotify import spotify_proxy/status import spotify_proxy/util @@ -29,9 +30,17 @@ pub fn start(_app, _type) -> Result(process.Pid, actor.StartError) { // // gg? + let ratelimit_name = process.new_name("ratelimiter") + let ratelimit_subject = process.named_subject(ratelimit_name) + let status_name = process.new_name("spotify_status") let status_subject = process.named_subject(status_name) + let ratelimit_child = + supervision.worker(fn() { + Ok(actor.Started(ratelimiter.spawn_link(ratelimit_name, 2, 10_000), Nil)) + }) + let spotify_child = supervision.worker(fn() { spotify.spawn_link( @@ -47,10 +56,11 @@ pub fn start(_app, _type) -> Result(process.Pid, actor.StartError) { status.spawn_link(status_name) |> result.map(actor.Started(_, Nil)) }) - let api_child = api.supervised(status_subject) + let api_child = api.supervised(status_subject, ratelimit_subject) let res = sup.new(sup.OneForOne) + |> sup.add(ratelimit_child) |> sup.add(spotify_child) |> sup.add(status_child) |> sup.add(api_child) diff --git a/spotify_proxy/src/spotify_proxy/api.gleam b/spotify_proxy/src/spotify_proxy/api.gleam index d6db3d9..a162b57 100644 --- a/spotify_proxy/src/spotify_proxy/api.gleam +++ b/spotify_proxy/src/spotify_proxy/api.gleam @@ -1,30 +1,54 @@ import gleam/erlang/process import gleam/int import gleam/json +import gleam/list import gleam/option.{type Option, None, Some} import gleam/otp/actor +import gleam/result +import gleam/string import mist +import spotify_proxy/ratelimiter import spotify_proxy/status import spotify_proxy/util import wisp import wisp/wisp_mist type Context { - Context(status_subject: process.Subject(status.Msg)) + Context( + status_subject: process.Subject(status.Msg), + ratelimit_subject: process.Subject(ratelimiter.Msg), + ) } -pub fn supervised(status_subject: process.Subject(status.Msg)) { +pub fn supervised( + status_subject: process.Subject(status.Msg), + ratelimit_subject: process.Subject(ratelimiter.Msg), +) { let secret_key_base = wisp.random_string(64) - let ctx = Context(status_subject:) + let ctx = Context(status_subject:, ratelimit_subject:) - wisp_mist.handler(fn(req) { handle_request(req, ctx) }, secret_key_base) + wisp_mist.handler(fn(req) { router(req, ctx) }, secret_key_base) |> mist.new |> mist.port(8000) |> mist.supervised } -fn handle_request(req: wisp.Request, ctx: Context) -> wisp.Response { +fn middleware( + req: wisp.Request, + ctx: Context, + handle_request: fn(wisp.Request) -> wisp.Response, +) -> wisp.Response { + use <- wisp.log_request(req) + use <- wisp.rescue_crashes + use req <- wisp.handle_head(req) + use req <- wisp.csrf_known_header_protection(req) + use <- ratelimit(req, ctx) + handle_request(req) +} + +fn router(req: wisp.Request, ctx: Context) -> wisp.Response { + use req <- middleware(req, ctx) case wisp.path_segments(req) { ["now-playing"] -> now_playing(req, ctx) _ -> wisp.not_found() @@ -77,3 +101,23 @@ fn artist_to_json(artist: status.Artist) -> json.Json { #("name", json.string(name)), ]) } + +fn ratelimit( + req: wisp.Request, + ctx: Context, + next: fn() -> wisp.Response, +) -> wisp.Response { + let ip = + req.headers + |> list.find_map(fn(header) { + case string.lowercase(header.0) { + h if h == "x-forwarded-for" -> Ok(header.1) + _ -> Error(Nil) + } + }) + |> result.unwrap(req.host) + case actor.call(ctx.ratelimit_subject, 200, ratelimiter.Take(_, ip)) { + True -> next() + False -> wisp.response(429) + } +} diff --git a/spotify_proxy/src/spotify_proxy/ratelimiter.gleam b/spotify_proxy/src/spotify_proxy/ratelimiter.gleam new file mode 100644 index 0000000..f5eacec --- /dev/null +++ b/spotify_proxy/src/spotify_proxy/ratelimiter.gleam @@ -0,0 +1,146 @@ +import gleam/dict +import gleam/erlang/process +import gleam/int +import gleam/otp/actor +import gleam/result +import logging + +type State { + State( + ratelimits: dict.Dict(String, Ratelimit), + max: Int, + leak_every: Int, + self_subject: process.Subject(Msg), + ) +} + +type Ratelimit { + Ratelimit(count: Int, max: Int, leak_every: Int) +} + +pub type Msg { + Take(process.Subject(Bool), String) + Leak(String) + Log +} + +pub fn spawn_link(name: process.Name(Msg), max: Int, leak_every: Int) { + let self_subject = process.named_subject(name) + + let init = fn(_sub) { + let selector = process.new_selector() |> process.select(self_subject) + let state = State(ratelimits: dict.new(), max:, leak_every:, self_subject:) + actor.initialised(state) + |> actor.selecting(selector) + |> actor.returning(Nil) + } + + let assert Ok(actor.Started(pid, _data)) = + actor.new_with_initialiser(100, fn(sub) { Ok(init(sub)) }) + |> actor.on_message(handle_message) + |> actor.start + + let _ = process.register(pid, name) + process.send(self_subject, Log) + + pid +} + +fn handle_message(state: State, msg: Msg) -> actor.Next(State, Msg) { + let state = case msg { + Log -> { + logging.log( + logging.Info, + "Ratelimiter is currently tracking " + <> int.to_string(dict.size(state.ratelimits)) + <> " ratelimits", + ) + + process.send_after(state.self_subject, 5 * 60 * 1000, Log) + state + } + // SHOULD BE QUEUED AFTER EVERY TAKE! this will eventually equate to 0! + Leak(key) -> { + case dict.get(state.ratelimits, key) { + Ok(ratelimit) -> { + // remove the ratelimit altogether if it's empty for mem usage concerns + // (does this even make a difference?) + let ratelimits = case leak(ratelimit) { + ratelimit if ratelimit.count == 0 -> { + logging.log( + logging.Debug, + "ratelimit for " <> key <> " empty, removing...", + ) + dict.delete(state.ratelimits, key) + } + ratelimit -> dict.insert(state.ratelimits, key, ratelimit) + } + + State(..state, ratelimits:) + } + Error(Nil) -> { + logging.log( + logging.Debug, + "Ratelimit not found to leak; it was already empty!", + ) + state + } + } + } + Take(return, key) -> { + let rl = + state.ratelimits + |> dict.get(key) + |> result.unwrap(Ratelimit( + count: 0, + max: state.max, + leak_every: state.leak_every, + )) + |> take + + let rl = case rl { + Ok(rl) -> { + // successfully taken + process.send(return, True) + rl + } + Error(rl) -> { + // ratelimited + process.send(return, False) + rl + } + } + + let ratelimits = dict.insert(state.ratelimits, key, rl) + process.send_after(state.self_subject, state.leak_every, Leak(key)) + State(..state, ratelimits:) + } + } + + actor.continue(state) +} + +fn leak(ratelimit: Ratelimit) -> Ratelimit { + logging.log(logging.Debug, "Leaking") + let ratelimit = Ratelimit(..ratelimit, count: int.max(0, ratelimit.count - 1)) + logging.log(logging.Debug, "new count: " <> int.to_string(ratelimit.count)) + ratelimit +} + +/// Ok(rl) if successfully taken, Error(rl) if ratelimited +fn take(ratelimit: Ratelimit) -> Result(Ratelimit, Ratelimit) { + logging.log(logging.Debug, "Taking") + case ratelimit.count { + count if count == ratelimit.max -> { + Error(ratelimit) + } + count -> { + let ratelimit = Ratelimit(..ratelimit, count: count + 1) + logging.log( + logging.Debug, + "new count: " <> int.to_string(ratelimit.count), + ) + Ok(ratelimit) + } + } +}