diff --git a/README.md b/README.md index 06ff110..d4e7f6d 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,8 @@ # Latch -**TODO: Add description** +atproto OAuth library attempting to follow the specification strictly, while also following Elixir library guidelines. The goal is for the library to be easy to use and not get in your way, but fully flexible. + +This is a pretty extensive introduction to atproto OAuth [Beyond the Statusphere: Part 2, ATProto OAuth, the TLDR](https://leaflet.pub/77df80c7-ec7e-4728-afa9-e367d99adb97). The core of this code originates from [annot.at](https://annot.at). ## Installation @@ -18,3 +20,16 @@ end Documentation can be generated with [ExDoc](https://github.com/elixir-lang/ex_doc) and published on [HexDocs](https://hexdocs.pm). Once published, the docs can be found at . + +## Roadmap + +- [x] Confidential client +- [x] DPoP nonce caching +- [ ] Public client +- [ ] Local client +- [ ] Extensive tests + +## Specification references + +* https://docs.bsky.app/docs/advanced-guides/oauth-client +* https://atproto.com/specs/oauth diff --git a/lib/latch.ex b/lib/latch.ex index 814c0aa..e16bf77 100644 --- a/lib/latch.ex +++ b/lib/latch.ex @@ -3,6 +3,8 @@ defmodule Latch do A library for building atproto OAuth integrations. Stateless. """ + use Supervisor + alias Latch.Config alias Latch.Discovery alias Latch.DPoP @@ -16,7 +18,41 @@ defmodule Latch do alias Latch.Request alias Latch.Session - def client_metadata(%Config{} = config) do + def start_link(opts) do + name = Keyword.fetch!(opts, :name) + Supervisor.start_link(__MODULE__, opts, name: name) + end + + def child_spec(opts) do + name = Keyword.fetch!(opts, :name) + + %{ + id: name, + type: :supervisor, + start: {__MODULE__, :start_link, [opts]} + } + end + + @impl Supervisor + def init(opts) do + config = Keyword.fetch!(opts, :config) + name = Keyword.fetch!(opts, :name) + config = Map.put(config, :name, name) + + :persistent_term.put({__MODULE__, self()}, config) + :persistent_term.put({__MODULE__, name}, config) + + Supervisor.init( + [ + {Latch.NonceCache, opts} + ], + strategy: :one_for_one + ) + end + + def client_metadata(name) do + %Config{} = config = config(name) + Latch.ClientMetadata.build( client_id: config.client_id, redirect_uris: [config.redirect_uri], @@ -27,15 +63,18 @@ defmodule Latch do ) end - def authorize(handle, %Config{} = config) do + def authorize(name, handle) do + %Config{} = config = config(name) + verifier = PKCE.generate_verifier() dpop_key = DPoP.generate_key() state = Base.url_encode64(:crypto.strong_rand_bytes(32), padding: false) + session_id = Base.url_encode64(:crypto.strong_rand_bytes(32), padding: false) with {:ok, identity} <- Identity.resolve_handle(handle), {:ok, server} <- Discovery.discover(identity.pds_endpoint), {:ok, request_uri} <- - Flow.par(server, + Flow.par(config, session_id, server, client_id: config.client_id, client_jwk: config.signing_key, redirect_uri: config.redirect_uri, @@ -53,27 +92,32 @@ defmodule Latch do issuer: server.issuer, token_endpoint: server.token_endpoint, pkce_verifier: verifier, - dpop_key: dpop_key + dpop_key: dpop_key, + session_id: session_id }, :ok <- store_request(config, request) do {:ok, Flow.authorization_url(server, config.client_id, request_uri)} end end - def callback(%{"state" => state} = params, %Config{} = config) when is_binary(state) do + def callback(name, %{"state" => state} = params) when is_binary(state) do + %Config{} = config = config(name) + with {:ok, request} <- take_request(config, state), :ok <- verify_state(request, state) do complete_callback(params, request, config) end end - def callback(_params, %Config{}) do + def callback(_name, _params) do {:error, %InvalidResponse{reason: :unexpected_response}} end - def refresh(%Session{} = session, %Config{} = config) do + def refresh(name, %Session{} = session) do + %Config{} = config = config(name) + with {:ok, server} <- Discovery.discover(session.pds_endpoint) do - Flow.refresh(server, session, + Flow.refresh(config, session.session_id, server, session, client_id: config.client_id, client_jwk: config.signing_key ) @@ -98,12 +142,12 @@ defmodule Latch do defp complete_callback( %{"code" => code, "iss" => issuer}, - request, + %Request{} = request, config ) when is_binary(code) do with :ok <- verify_issuer(request, issuer) do - Flow.exchange_code( + Flow.exchange_code(config, request.session_id, client_id: config.client_id, client_jwk: config.signing_key, redirect_uri: config.redirect_uri, @@ -113,7 +157,8 @@ defmodule Latch do expected_did: request.did, pds_endpoint: request.pds_endpoint, issuer: request.issuer, - token_endpoint: request.token_endpoint + token_endpoint: request.token_endpoint, + session_id: request.session_id ) end end @@ -157,4 +202,8 @@ defmodule Latch do defp verify_issuer(_request, _issuer) do {:error, %SecurityViolation{reason: :issuer_mismatch}} end + + defp config(name) do + :persistent_term.get({__MODULE__, name}) + end end diff --git a/lib/latch/client.ex b/lib/latch/client.ex index d5cf260..810c186 100644 --- a/lib/latch/client.ex +++ b/lib/latch/client.ex @@ -25,7 +25,7 @@ defmodule Latch.Client do @spec query(Config.t(), String.t(), String.t(), keyword()) :: {:ok, map()} | {:error, Error.t()} def query(%Config{} = config, did, method, params \\ []) do - call(config, did, fn session -> XRPC.query(session, method, params) end) + call(config, did, fn session -> XRPC.query(config, session, method, params) end) end @doc """ @@ -34,7 +34,7 @@ defmodule Latch.Client do @spec procedure(Config.t(), String.t(), String.t(), map()) :: {:ok, map()} | {:error, Error.t()} def procedure(%Config{} = config, did, method, body) do - call(config, did, fn session -> XRPC.procedure(session, method, body) end) + call(config, did, fn session -> XRPC.procedure(config, session, method, body) end) end @doc """ @@ -43,7 +43,7 @@ defmodule Latch.Client do @spec upload_blob(Config.t(), String.t(), binary(), String.t()) :: {:ok, map()} | {:error, Error.t()} def upload_blob(%Config{} = config, did, bytes, content_type) do - call(config, did, fn session -> XRPC.upload_blob(session, bytes, content_type) end) + call(config, did, fn session -> XRPC.upload_blob(config, session, bytes, content_type) end) end defp call(config, did, fun) do @@ -100,7 +100,7 @@ defmodule Latch.Client do defp do_refresh(config, session) do result = with {:ok, server} <- Discovery.discover(session.pds_endpoint) do - Flow.refresh(server, session, + Flow.refresh(config, session.session_id, server, session, client_id: config.client_id, client_jwk: config.signing_key ) diff --git a/lib/latch/config.ex b/lib/latch/config.ex index da0d3cb..b0e5dc4 100644 --- a/lib/latch/config.ex +++ b/lib/latch/config.ex @@ -15,7 +15,7 @@ defmodule Latch.Config do @default_request_ttl 600 @enforce_keys [:store, :client_id, :redirect_uri, :scope, :signing_key] - defstruct @enforce_keys ++ [:client_name, :client_uri, request_ttl: @default_request_ttl] + defstruct @enforce_keys ++ [:client_name, :client_uri, :name, request_ttl: @default_request_ttl] @type t :: %__MODULE__{ store: module(), @@ -24,6 +24,7 @@ defmodule Latch.Config do scope: String.t(), signing_key: JOSE.JWK.t(), client_name: String.t() | nil, - client_uri: String.t() | nil + client_uri: String.t() | nil, + name: term() } end diff --git a/lib/latch/flow.ex b/lib/latch/flow.ex index 0b00d00..14e4982 100644 --- a/lib/latch/flow.ex +++ b/lib/latch/flow.ex @@ -11,6 +11,7 @@ defmodule Latch.Flow do """ alias Latch.ClientAssertion + alias Latch.Config alias Latch.DPoP alias Latch.Error.InvalidResponse alias Latch.Error.MissingDPoPNonce @@ -34,10 +35,10 @@ defmodule Latch.Flow do - `:login_hint` - the user's handle or DID """ - @spec par(ServerMetadata.t(), keyword()) :: + @spec par(Config.t(), String.t(), ServerMetadata.t(), keyword()) :: {:ok, String.t()} | {:error, InvalidResponse.t() | MissingDPoPNonce.t() | OAuth.t() | Transport.t()} - def par(%ServerMetadata{} = server, opts) do + def par(%Config{} = config, session_id, %ServerMetadata{} = server, opts) do client_id = Keyword.fetch!(opts, :client_id) client_jwk = Keyword.fetch!(opts, :client_jwk) redirect_uri = Keyword.fetch!(opts, :redirect_uri) @@ -65,7 +66,8 @@ defmodule Latch.Flow do ) end - with {:ok, body} <- dpop_request(server.par_endpoint, build_form, dpop_key) do + with {:ok, body} <- + dpop_request(config, session_id, server.par_endpoint, build_form, dpop_key) do parse_par_response(body) end end @@ -100,7 +102,7 @@ defmodule Latch.Flow do ## Optional - `:now` - base time for `expires_at` (defaults to the current time) """ - @spec exchange_code(keyword()) :: + @spec exchange_code(Config.t(), String.t(), keyword()) :: {:ok, Session.t()} | {:error, InvalidResponse.t() @@ -108,7 +110,7 @@ defmodule Latch.Flow do | OAuth.t() | SecurityViolation.t() | Transport.t()} - def exchange_code(opts) do + def exchange_code(%Config{} = config, session_id, opts) do client_id = Keyword.fetch!(opts, :client_id) client_jwk = Keyword.fetch!(opts, :client_jwk) redirect_uri = Keyword.fetch!(opts, :redirect_uri) @@ -133,10 +135,10 @@ defmodule Latch.Flow do ] end - with {:ok, body} <- dpop_request(token_endpoint, build_form, dpop_key), + with {:ok, body} <- dpop_request(config, session_id, token_endpoint, build_form, dpop_key), {:ok, tokens} <- parse_token_response(body), :ok <- verify_sub(tokens.sub, expected_did) do - {:ok, build_session(tokens, issuer, pds_endpoint, dpop_key, now)} + {:ok, build_session(tokens, issuer, pds_endpoint, dpop_key, session_id, now)} end end @@ -154,7 +156,7 @@ defmodule Latch.Flow do ## Optional - `:now` - base time for `expires_at` (defaults to the current time) """ - @spec refresh(ServerMetadata.t(), Session.t(), keyword()) :: + @spec refresh(Config.t(), String.t(), ServerMetadata.t(), Session.t(), keyword()) :: {:ok, Session.t()} | {:error, InvalidResponse.t() @@ -162,7 +164,7 @@ defmodule Latch.Flow do | OAuth.t() | SecurityViolation.t() | Transport.t()} - def refresh(%ServerMetadata{} = server, %Session{} = session, opts) do + def refresh(config, session_id, %ServerMetadata{} = server, %Session{} = session, opts) do client_id = Keyword.fetch!(opts, :client_id) client_jwk = Keyword.fetch!(opts, :client_jwk) now = Keyword.get_lazy(opts, :now, &DateTime.utc_now/0) @@ -179,25 +181,49 @@ defmodule Latch.Flow do with :ok <- verify_refresh_issuer(server.issuer, session.issuer), {:ok, body} <- - dpop_request(server.token_endpoint, build_form, session.dpop_key), + dpop_request(config, session_id, server.token_endpoint, build_form, session.dpop_key), {:ok, tokens} <- parse_token_response(body), :ok <- verify_sub(tokens.sub, session.did) do - {:ok, build_session(tokens, server.issuer, session.pds_endpoint, session.dpop_key, now)} + {:ok, + build_session( + tokens, + server.issuer, + session.pds_endpoint, + session.dpop_key, + session.session_id, + now + )} end end - defp dpop_request(url, build_form, dpop_key, nonce \\ nil) do + defp dpop_request(config, session_id, url, build_form, dpop_key) do + origin = origin(url) + + nonce = + case Latch.NonceCache.get_nonce(config, session_id, origin) do + {:ok, nonce} -> nonce + :error -> nil + end + + send_dpop(config, session_id, url, build_form, dpop_key, origin, nonce) + end + + defp send_dpop(config, session_id, url, build_form, dpop_key, origin, nonce) do proof = DPoP.proof(dpop_key, "POST", url, nonce: nonce) with {:ok, %{status: status, body: raw, headers: headers}} <- HTTP.post_form(url, build_form.(), [{"dpop", proof}]), {:ok, body} <- decode_json(raw) do + if fresh = nonce_header(headers) do + Latch.NonceCache.put_nonce(config, session_id, origin, fresh) + end + cond do status in 200..299 -> {:ok, body} retry_nonce?(body, nonce) -> - retry_with_nonce(url, build_form, dpop_key, headers) + retry_with_nonce(config, session_id, url, build_form, dpop_key, origin, headers) is_binary(Map.get(body, "error")) -> {:error, @@ -213,14 +239,25 @@ defmodule Latch.Flow do end end - defp retry_with_nonce(url, build_form, dpop_key, headers) do + defp retry_with_nonce(config, session_id, url, build_form, dpop_key, origin, headers) do if nonce = nonce_header(headers) do - dpop_request(url, build_form, dpop_key, nonce) + send_dpop(config, session_id, url, build_form, dpop_key, origin, nonce) else {:error, %MissingDPoPNonce{}} end end + defp origin(url) do + %URI{scheme: scheme, host: host, port: port} = URI.parse(url) + default = if scheme == "https", do: 443, else: 80 + + if is_nil(port) or port == default do + "#{scheme}://#{host}" + else + "#{scheme}://#{host}:#{port}" + end + end + defp retry_nonce?(%{"error" => "use_dpop_nonce"}, nil), do: true defp retry_nonce?(_body, _nonce), do: false @@ -270,7 +307,7 @@ defmodule Latch.Flow do defp verify_sub(sub, sub), do: :ok defp verify_sub(_sub, _expected), do: {:error, %SecurityViolation{reason: :did_mismatch}} - defp build_session(tokens, issuer, pds_endpoint, dpop_key, now) do + defp build_session(tokens, issuer, pds_endpoint, dpop_key, session_id, now) do %Session{ did: tokens.sub, access_token: tokens.access_token, @@ -279,7 +316,8 @@ defmodule Latch.Flow do scope: tokens.scope, issuer: issuer, pds_endpoint: pds_endpoint, - expires_at: DateTime.add(now, tokens.expires_in, :second) + expires_at: DateTime.add(now, tokens.expires_in, :second), + session_id: session_id } end end diff --git a/lib/latch/nonce_cache.ex b/lib/latch/nonce_cache.ex new file mode 100644 index 0000000..d3712cc --- /dev/null +++ b/lib/latch/nonce_cache.ex @@ -0,0 +1,92 @@ +defmodule Latch.NonceCache do + @moduledoc false + + use GenServer + + @table_options [ + :set, + :public, + :named_table, + read_concurrency: true, + write_concurrency: true + ] + + @default_sweep_interval :timer.minutes(2) + @default_ttl_ms :timer.minutes(5) + + def get_nonce(config, session_id, origin) do + key = {session_id, origin} + + case :ets.lookup(table(config), key) do + [{^key, nonce, expires_at}] -> + if System.monotonic_time(:millisecond) < expires_at do + {:ok, nonce} + else + :error + end + + [] -> + :error + end + end + + def put_nonce(config, session_id, origin, nonce, ttl_ms \\ @default_ttl_ms) do + expires_at = System.monotonic_time(:millisecond) + ttl_ms + :ets.insert(table(config), {{session_id, origin}, nonce, expires_at}) + :ok + end + + def row_count(config) do + :ets.info(table(config), :size) + end + + def start_link(opts) do + GenServer.start_link(__MODULE__, opts) + end + + @impl GenServer + def init(opts) do + config = Keyword.fetch!(opts, :config) + name = Keyword.fetch!(opts, :name) + table_name = :"latch_#{name}_nonce_cache" + table = :ets.new(table_name, @table_options) + + sweep_after = Keyword.get(opts, :sweep_after, @default_sweep_interval) + sweep_disabled = Keyword.get(opts, :sweep_disabled, false) + + schedule_sweep(sweep_after, sweep_disabled) + + {:ok, + %{ + config: config, + table: table, + sweep_after: sweep_after, + sweep_disabled: sweep_disabled + }} + end + + @impl GenServer + def handle_info(:sweep, state) do + sweep(table(state.config)) + schedule_sweep(state.sweep_after, state.sweep_disabled) + {:noreply, state} + end + + defp schedule_sweep(sweep_after, sweep_disabled) do + if not sweep_disabled do + Process.send_after(self(), :sweep, sweep_after) + end + end + + defp sweep(table) do + now = System.monotonic_time(:millisecond) + + :ets.select_delete(table, [ + {{:"$1", :"$2", :"$3"}, [{:<, :"$3", now}], [true]} + ]) + end + + defp table(config) do + :"latch_#{config.name}_nonce_cache" + end +end diff --git a/lib/latch/request.ex b/lib/latch/request.ex index 3d1ae58..27b74b7 100644 --- a/lib/latch/request.ex +++ b/lib/latch/request.ex @@ -9,7 +9,8 @@ defmodule Latch.Request do :issuer, :token_endpoint, :pkce_verifier, - :dpop_key + :dpop_key, + :session_id ] defstruct @enforce_keys @@ -21,6 +22,7 @@ defmodule Latch.Request do issuer: String.t(), token_endpoint: String.t(), pkce_verifier: String.t(), - dpop_key: JOSE.JWK.t() + dpop_key: JOSE.JWK.t(), + session_id: String.t() } end diff --git a/lib/latch/session.ex b/lib/latch/session.ex index 4bb72bd..bbe27d9 100644 --- a/lib/latch/session.ex +++ b/lib/latch/session.ex @@ -17,7 +17,8 @@ defmodule Latch.Session do :scope, :issuer, :pds_endpoint, - :expires_at + :expires_at, + :session_id ] defstruct @enforce_keys @@ -29,6 +30,7 @@ defmodule Latch.Session do scope: String.t(), issuer: String.t(), pds_endpoint: String.t(), - expires_at: DateTime.t() + expires_at: DateTime.t(), + session_id: String.t() } end diff --git a/lib/latch/xrpc.ex b/lib/latch/xrpc.ex index 3ddf4b4..9c3c864 100644 --- a/lib/latch/xrpc.ex +++ b/lib/latch/xrpc.ex @@ -8,6 +8,7 @@ defmodule Latch.XRPC do persistence are the caller's concerns, not this module's. """ + alias Latch.Config alias Latch.DPoP alias Latch.Error.InvalidResponse alias Latch.Error.MissingDPoPNonce @@ -22,9 +23,10 @@ defmodule Latch.XRPC do @doc """ Performs and authenticated XRPC query against the session's PDS. """ - @spec query(Session.t(), String.t(), keyword()) :: {:ok, map()} | {:error, error()} - def query(%Session{} = session, method, params \\ []) do + @spec query(Config.t(), Session.t(), String.t(), keyword()) :: {:ok, map()} | {:error, error()} + def query(%Config{} = config, %Session{} = session, method, params \\ []) do request( + config, session, "GET", session.pds_endpoint <> "/xrpc/" <> method <> query_string(params), @@ -35,18 +37,20 @@ defmodule Latch.XRPC do @doc """ Performs and authenticated XRPC procedure against the session's PDS. """ - @spec procedure(Session.t(), String.t(), map()) :: {:ok, map()} | {:error, error()} - def procedure(%Session{} = session, method, body) do - request(session, "POST", session.pds_endpoint <> "/xrpc/" <> method, {:json, body}) + @spec procedure(Config.t(), Session.t(), String.t(), map()) :: {:ok, map()} | {:error, error()} + def procedure(%Config{} = config, %Session{} = session, method, body) do + request(config, session, "POST", session.pds_endpoint <> "/xrpc/" <> method, {:json, body}) end @doc """ Uploads raw bytes of content_type as a blob, returning the response with the blog reference. """ - @spec upload_blob(Session.t(), binary(), String.t()) :: {:ok, map()} | {:error, error()} - def upload_blob(%Session{} = session, bytes, content_type) do + @spec upload_blob(Config.t(), Session.t(), binary(), String.t()) :: + {:ok, map()} | {:error, error()} + def upload_blob(%Config{} = config, %Session{} = session, bytes, content_type) do request( + config, session, "POST", session.pds_endpoint <> "/xrpc/com.atproto.repo.uploadBlob", @@ -54,7 +58,19 @@ defmodule Latch.XRPC do ) end - defp request(%Session{} = session, http_method, url, body, nonce \\ nil) do + defp request(%Config{} = config, %Session{} = session, http_method, url, body) do + origin = origin(url) + + nonce = + case Latch.NonceCache.get_nonce(config, session.session_id, origin) do + {:ok, nonce} -> nonce + :error -> nil + end + + send_dpop(config, session, http_method, url, body, origin, nonce) + end + + defp send_dpop(%Config{} = config, %Session{} = session, http_method, url, body, origin, nonce) do proof = DPoP.proof(session.dpop_key, http_method, url, nonce: nonce, @@ -65,8 +81,12 @@ defmodule Latch.XRPC do with {:ok, %{status: status, body: raw, headers: resp_headers}} <- HTTP.request(http_method, url, headers, body) do - if is_nil(nonce) and needs_nonce?(status, resp_headers) do - retry_with_nonce(session, http_method, url, body, resp_headers) + if fresh = nonce_header(resp_headers) do + Latch.NonceCache.put_nonce(config, session.session_id, origin, fresh) + end + + if needs_nonce?(status, resp_headers) do + retry_with_nonce(config, session, http_method, url, body, origin, resp_headers) else handle_response(status, raw) end @@ -83,10 +103,30 @@ defmodule Latch.XRPC do end end - defp retry_with_nonce(session, http_method, url, body, headers) do - case nonce_header(headers) do - nil -> {:error, %MissingDPoPNonce{}} - nonce -> request(session, http_method, url, body, nonce) + defp retry_with_nonce( + %Config{} = config, + %Session{} = session, + http_method, + url, + body, + origin, + headers + ) do + if nonce = nonce_header(headers) do + send_dpop(config, session, http_method, url, body, origin, nonce) + else + {:error, %MissingDPoPNonce{}} + end + end + + defp origin(url) do + %URI{scheme: scheme, host: host, port: port} = URI.parse(url) + default = if scheme == "https", do: 443, else: 80 + + if is_nil(port) or port == default do + "#{scheme}://#{host}" + else + "#{scheme}://#{host}:#{port}" end end diff --git a/test/latch/client_test.exs b/test/latch/client_test.exs index f726a6b..1112f3f 100644 --- a/test/latch/client_test.exs +++ b/test/latch/client_test.exs @@ -35,13 +35,16 @@ defmodule Latch.ClientTest do expect(Discovery, :discover, fn @pds -> {:ok, server} end) - expect(Flow, :refresh, fn ^server, ^stale_session, opts -> + expect(Flow, :refresh, fn _config, _session_id, ^server, ^stale_session, opts -> assert opts[:client_id] == config.client_id assert opts[:client_jwk] == config.signing_key {:ok, refreshed_session} end) - expect(XRPC, :query, fn ^refreshed_session, "app.bsky.actor.getProfile", actor: @did -> + expect(XRPC, :query, fn _config, + ^refreshed_session, + "app.bsky.actor.getProfile", + actor: @did -> {:ok, %{"did" => @did}} end) @@ -67,16 +70,16 @@ defmodule Latch.ClientTest do :ok = Latch.TestStore.put_session(@did, stale_session) expect(XRPC, :query, 2, fn - ^stale_session, "app.bsky.actor.getProfile", actor: @did -> + _config, ^stale_session, "app.bsky.actor.getProfile", actor: @did -> {:error, %XRPCError{status: 401, body: %{}}} - ^refreshed_session, "app.bsky.actor.getProfile", actor: @did -> + _config, ^refreshed_session, "app.bsky.actor.getProfile", actor: @did -> {:ok, %{"did" => @did}} end) expect(Discovery, :discover, fn @pds -> {:ok, server} end) - expect(Flow, :refresh, fn ^server, ^stale_session, _opts -> + expect(Flow, :refresh, fn _config, _session_id, ^server, ^stale_session, _opts -> {:ok, refreshed_session} end) @@ -96,7 +99,8 @@ defmodule Latch.ClientTest do scope: "atproto", issuer: @issuer, pds_endpoint: @pds, - expires_at: expires_at + expires_at: expires_at, + session_id: "random 32 chars" } end diff --git a/test/latch/flow_test.exs b/test/latch/flow_test.exs index d83d224..037737d 100644 --- a/test/latch/flow_test.exs +++ b/test/latch/flow_test.exs @@ -14,6 +14,21 @@ defmodule Latch.FlowTest do access_token = "access-token" refresh_token = "refresh-token" + config = %Latch.Config{ + store: Latch.TestStore, + client_id: "https://client.example.com/oauth-client-metadata.json", + redirect_uri: "https://client.example.com/oauth/callback", + scope: "atproto", + signing_key: nil, + name: :"flow_test_#{inspect(self())}" + } + + start_link_supervised!( + {Latch.NonceCache, config: config, name: config.name, sweep_disabled: true} + ) + + session_id = "session-id" + client_jwk = JOSE.JWK.generate_key({:ec, "P-256"}) dpop_key = JOSE.JWK.generate_key({:ec, "P-256"}) @@ -43,7 +58,7 @@ defmodule Latch.FlowTest do end) assert {:ok, session} = - Flow.exchange_code( + Flow.exchange_code(config, session_id, client_id: "https://client.example.com/oauth-client-metadata.json", client_jwk: client_jwk, redirect_uri: "https://client.example.com/oauth/callback", @@ -54,7 +69,8 @@ defmodule Latch.FlowTest do pds_endpoint: "https://pds.example.com", issuer: "https://issuer.example.com", token_endpoint: "https://issuer.example.com/oauth/token", - now: ~U[2026-01-01 00:00:00Z] + now: ~U[2026-01-01 00:00:00Z], + session_id: "random 32 chars" ) assert session.did == did @@ -70,6 +86,21 @@ defmodule Latch.FlowTest do client_jwk = JOSE.JWK.generate_key({:ec, "P-256"}) dpop_key = JOSE.JWK.generate_key({:ec, "P-256"}) + config = %Latch.Config{ + store: Latch.TestStore, + client_id: "https://client.example.com/oauth-client-metadata.json", + redirect_uri: "https://client.example.com/oauth/callback", + scope: "atproto", + signing_key: nil, + name: :"flow_test_#{inspect(self())}" + } + + start_link_supervised!( + {Latch.NonceCache, config: config, name: config.name, sweep_disabled: true} + ) + + session_id = "session-id" + server = %ServerMetadata{ issuer: "https://issuer.example.com", authorization_endpoint: "https://issuer.example.com/oauth/authorize", @@ -95,7 +126,7 @@ defmodule Latch.FlowTest do end) assert {:ok, "urn:ietf:params:oauth:request_uri:request"} = - Flow.par(server, + Flow.par(config, session_id, server, client_id: "https://client.example.com/oauth-client-metadata.json", client_jwk: client_jwk, redirect_uri: "https://client.example.com/oauth/callback", @@ -112,6 +143,21 @@ defmodule Latch.FlowTest do test "rejects a refresh when discovery returns a different issuer" do reject(HTTP, :post_form, 3) + config = %Latch.Config{ + store: Latch.TestStore, + client_id: "https://client.example.com/oauth-client-metadata.json", + redirect_uri: "https://client.example.com/oauth/callback", + scope: "atproto", + signing_key: nil, + name: :"flow_test_#{inspect(self())}" + } + + start_link_supervised!( + {Latch.NonceCache, config: config, name: config.name, sweep_disabled: true} + ) + + session_id = "session-id" + server = %ServerMetadata{ issuer: "https://other.example.com", authorization_endpoint: "https://other.example.com/oauth/authorize", @@ -128,11 +174,15 @@ defmodule Latch.FlowTest do scope: "atproto", issuer: "https://issuer.example.com", pds_endpoint: "https://pds.example.com", - expires_at: ~U[2026-01-01 00:00:00Z] + expires_at: ~U[2026-01-01 00:00:00Z], + session_id: "random 32 chars" } assert {:error, %SecurityViolation{reason: :issuer_mismatch}} = - Flow.refresh(server, session, client_id: "client-id", client_jwk: nil) + Flow.refresh(config, session_id, server, session, + client_id: "client-id", + client_jwk: nil + ) end end end diff --git a/test/latch/nonce_cache_test.exs b/test/latch/nonce_cache_test.exs new file mode 100644 index 0000000..34aae0d --- /dev/null +++ b/test/latch/nonce_cache_test.exs @@ -0,0 +1,87 @@ +defmodule Latch.NonceCacheTest do + use ExUnit.Case, async: true + + alias Latch.Config + alias Latch.NonceCache + + describe "get_nonce/3" do + test "returns error if no nonce exists" do + name = :"nonce_test_#{inspect(self())}" + config = config(name) + session_id = "1" + origin = "example.com" + + _pid = + start_link_supervised!({NonceCache, config: config, name: name, sweep_disabled: true}) + + assert :error = NonceCache.get_nonce(config, session_id, origin) + end + end + + describe "put_nonce/3" do + test "returns nonce if exists" do + name = :"nonce_test_#{inspect(self())}" + config = config(name) + session_id = "1" + origin = "example.com" + nonce = "nonce" + + _pid = + start_link_supervised!({NonceCache, config: config, name: name, sweep_disabled: true}) + + # Empty cache, return error + assert :error = NonceCache.get_nonce(config, session_id, origin) + + # Store in cache + assert :ok = NonceCache.put_nonce(config, session_id, origin, nonce) + + # Returns now + assert {:ok, ^nonce} = NonceCache.get_nonce(config, session_id, origin) + + # Other things don't exist + assert :error = NonceCache.get_nonce(config, "other", origin) + assert :error = NonceCache.get_nonce(config, session_id, "other") + end + end + + describe "GenServer" do + test "sweeps old records" do + name = :"nonce_test_#{inspect(self())}" + config = config(name) + session_id = "1" + origin = "example.com" + nonce = "nonce" + + pid = + start_link_supervised!( + {NonceCache, config: config, name: name, sweep_after: 0, sweep_disabled: true} + ) + + assert NonceCache.row_count(config) == 0 + + # Set negative TTL to force sweep + ttl_ms = -1 + + # Store in cache + assert :ok = NonceCache.put_nonce(config, session_id, origin, nonce, ttl_ms) + + assert NonceCache.row_count(config) == 1 + + send(pid, :sweep) + :sys.get_state(pid) + + assert NonceCache.row_count(config) == 0 + end + end + + defp config(name) do + %Config{ + store: Latch.TestStore, + client_id: "client-id", + redirect_uri: "redirect-uri", + scope: "atproto", + signing_key: nil, + name: name + } + end +end diff --git a/test/latch_test.exs b/test/latch_test.exs index 685b3ca..b4f1b60 100644 --- a/test/latch_test.exs +++ b/test/latch_test.exs @@ -30,13 +30,14 @@ defmodule LatchTest do signing_key: nil } + pid = start_latch(config) identity = %Identity{did: @did, handle: @handle, pds_endpoint: @pds} server = server() expect(Identity, :resolve_handle, fn @handle -> {:ok, identity} end) expect(Discovery, :discover, fn @pds -> {:ok, server} end) - expect(Flow, :par, fn ^server, opts -> + expect(Flow, :par, fn _config, _session_id, ^server, opts -> assert opts[:client_id] == config.client_id assert opts[:redirect_uri] == config.redirect_uri assert opts[:scope] == "atproto" @@ -48,7 +49,7 @@ defmodule LatchTest do {:ok, request_uri} end) - assert {:ok, redirect_url} = Latch.authorize(@handle, config) + assert {:ok, redirect_url} = Latch.authorize(pid, @handle) assert URI.parse(redirect_url).path == "/oauth/authorize" @@ -89,6 +90,8 @@ defmodule LatchTest do signing_key: nil } + pid = start_latch(config) + request = %Request{ state: state, did: @did, @@ -97,7 +100,8 @@ defmodule LatchTest do issuer: @issuer, token_endpoint: @issuer <> "/oauth/token", pkce_verifier: "pkce-verifier", - dpop_key: dpop_key + dpop_key: dpop_key, + session_id: "32-random-characters" } session = %Session{ @@ -108,12 +112,13 @@ defmodule LatchTest do scope: "atproto", issuer: @issuer, pds_endpoint: @pds, - expires_at: ~U[2026-01-01 01:00:00Z] + expires_at: ~U[2026-01-01 01:00:00Z], + session_id: "32 random chars" } :ok = Latch.TestStore.put_request(state, request, 600) - expect(Flow, :exchange_code, fn opts -> + expect(Flow, :exchange_code, fn _config, _session_id, opts -> assert opts[:code] == code assert opts[:code_verifier] == "pkce-verifier" assert opts[:dpop_key] == dpop_key @@ -126,16 +131,16 @@ defmodule LatchTest do assert {:ok, ^session} = Latch.callback( + pid, %{ "state" => state, "iss" => @issuer, "code" => code - }, - config + } ) assert {:error, %Latch.Error.SecurityViolation{reason: :state_mismatch}} = - Latch.callback(%{"state" => state, "iss" => @issuer, "code" => code}, config) + Latch.callback(pid, %{"state" => state, "iss" => @issuer, "code" => code}) end end @@ -149,6 +154,8 @@ defmodule LatchTest do signing_key: nil } + pid = start_latch(config) + session = %Session{ did: @did, access_token: "access-token", @@ -157,7 +164,8 @@ defmodule LatchTest do scope: "atproto", issuer: @issuer, pds_endpoint: @pds, - expires_at: ~U[2026-01-01 00:00:00Z] + expires_at: ~U[2026-01-01 00:00:00Z], + session_id: "32 random chars" } refreshed_session = %{session | access_token: "refreshed-access-token"} @@ -165,13 +173,13 @@ defmodule LatchTest do expect(Discovery, :discover, fn @pds -> {:ok, server} end) - expect(Flow, :refresh, fn ^server, ^session, opts -> + expect(Flow, :refresh, fn _config, _session_id, ^server, ^session, opts -> assert opts[:client_id] == config.client_id assert opts[:client_jwk] == config.signing_key {:ok, refreshed_session} end) - assert {:ok, ^refreshed_session} = Latch.refresh(session, config) + assert {:ok, ^refreshed_session} = Latch.refresh(pid, session) end end @@ -184,4 +192,9 @@ defmodule LatchTest do scopes_supported: ["atproto"] } end + + defp start_latch(config) do + name = String.to_atom("latch_#{inspect(self())}") + start_link_supervised!({Latch, name: name, config: config}) + end end