diff --git a/.gitignore b/.gitignore index 96b4619..f8f2ae8 100644 --- a/.gitignore +++ b/.gitignore @@ -6,6 +6,7 @@ erl_crash.dump *.ez atex-*.tar /tmp/ +/priv/dets/ .envrc .direnv diff --git a/AGENTS.md b/AGENTS.md index cf53021..3a03bf9 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -21,6 +21,7 @@ TypedStruct for structs - **Moduledocs**: All public modules need `@moduledoc`, public functions need `@doc` with examples + - When writing lists in documentation, use `-` as the list character. - **Error Handling**: Return `{:ok, result}` or `{:error, reason}` tuples; use pattern matching in case statements - **Pattern Matching**: Prefer pattern matching over conditionals; use guards diff --git a/CHANGELOG.md b/CHANGELOG.md index 45c2b9f..4b5705e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,11 +10,15 @@ and this project adheres to ### Breaking Changes -- Rename `Atex.XRPC.OAuthClient.update_plug/2` to `update_conn/2`, to match the - naming of `from_conn/1`. - `Atex.OAuth.Plug` now raises `Atex.OAuth.Error` exceptions instead of handling error situations internally. Applications should implement `Plug.ErrorHandler` to catch and gracefully handle them. +- `Atex.OAuth.Plug` now saves only the user's DID in the session instead of the + entire OAuth session object. Applications must use `Atex.OAuth.SessionStore` + to manage OAuth sessions. +- `Atex.XRPC.OAuthClient` has been overhauled to use `Atex.OAuth.SessionStore` + for retrieving and managing OAuth sessions, making it easier to use with not + needing to manually keep a Plug session in sync. ### Added diff --git a/lib/atex/xrpc/oauth_client.ex b/lib/atex/xrpc/oauth_client.ex index 7d0cb1d..239e77f 100644 --- a/lib/atex/xrpc/oauth_client.ex +++ b/lib/atex/xrpc/oauth_client.ex @@ -1,122 +1,151 @@ defmodule Atex.XRPC.OAuthClient do + @moduledoc """ + OAuth client for making authenticated XRPC requests to AT Protocol servers. + + The client contains a user's DID and talks to `Atex.OAuth.SessionStore` to + retrieve sessions internally to make requests. As a result, it will only work + for users that have gone through an OAuth flow; see `Atex.OAuth.Plug` for an + existing method of doing that. + + The entire OAuth session lifecycle is handled transparently, with the access + token being refreshed automatically as required. + + ## Usage + + # Create from an existing OAuth session + {:ok, client} = Atex.XRPC.OAuthClient.new("did:plc:abc123") + + # Or extract from a Plug.Conn after OAuth flow + {:ok, client} = Atex.XRPC.OAuthClient.from_conn(conn) + + # Make XRPC requests + {:ok, response, client} = Atex.XRPC.get(client, "com.atproto.repo.listRecords") + """ + alias Atex.OAuth - alias Atex.XRPC use TypedStruct @behaviour Atex.XRPC.Client typedstruct enforce: true do - field :endpoint, String.t() - field :issuer, String.t() - field :access_token, String.t() - field :refresh_token, String.t() field :did, String.t() - field :expires_at, NaiveDateTime.t() - field :dpop_nonce, String.t() | nil, enforce: false - field :dpop_key, JOSE.JWK.t() end @doc """ - Create a new OAuthClient struct. + Create a new OAuthClient from a DID. + + Validates that an OAuth session exists for the given DID in the session store + before returning the client struct. + + ## Examples + + iex> Atex.XRPC.OAuthClient.new("did:plc:abc123") + {:ok, %Atex.XRPC.OAuthClient{did: "did:plc:abc123"}} + + iex> Atex.XRPC.OAuthClient.new("did:plc:nosession") + {:error, :not_found} + """ - @spec new( - String.t(), - String.t(), - String.t(), - String.t(), - NaiveDateTime.t(), - JOSE.JWK.t(), - String.t() | nil - ) :: t() - def new(endpoint, did, access_token, refresh_token, expires_at, dpop_key, dpop_nonce) do - {:ok, issuer} = OAuth.get_authorization_server(endpoint) - - %__MODULE__{ - endpoint: endpoint, - issuer: issuer, - access_token: access_token, - refresh_token: refresh_token, - did: did, - expires_at: expires_at, - dpop_nonce: dpop_nonce, - dpop_key: dpop_key - } + @spec new(String.t()) :: {:ok, t()} | {:error, atom()} + def new(did) do + # Make sure session exists before returning a struct + case Atex.OAuth.SessionStore.get(did) do + {:ok, _session} -> + {:ok, %__MODULE__{did: did}} + + err -> + err + end end @doc """ - Create an OAuthClient struct from a `Plug.Conn`. + Create an OAuthClient from a `Plug.Conn`. + + Extracts the DID from the session (stored under `:atex_session` key) and validates + that the OAuth session is still valid. If the token is expired or expiring soon, + it attempts to refresh it. + + Requires the conn to have passed through `Plug.Session` and `Plug.Conn.fetch_session/2`. + + ## Returns - Requires the conn to have passed through `Plug.Session` and - `Plug.Conn.fetch_session/2` so that the session can be acquired and have the - `atex_oauth` key fetched from it. + - `{:ok, client}` - Successfully created client + - `{:error, :reauth}` - Session exists but refresh failed, user needs to re-authenticate + - `:error` - No session found in conn + + ## Examples + + # After OAuth flow completes + conn = Plug.Conn.put_session(conn, :atex_session, "did:plc:abc123") + {:ok, client} = Atex.XRPC.OAuthClient.from_conn(conn) - Returns `:error` if the state is missing or is not the expected shape. """ - @spec from_conn(Plug.Conn.t()) :: {:ok, t()} | :error + @spec from_conn(Plug.Conn.t()) :: {:ok, t()} | :error | {:error, atom()} def from_conn(%Plug.Conn{} = conn) do - oauth_state = Plug.Conn.get_session(conn, :atex_oauth) - - case oauth_state do - %{ - access_token: access_token, - refresh_token: refresh_token, - did: did, - pds: pds, - expires_at: expires_at, - dpop_nonce: dpop_nonce, - dpop_key: dpop_key - } -> - {:ok, new(pds, did, access_token, refresh_token, expires_at, dpop_key, dpop_nonce)} + oauth_did = Plug.Conn.get_session(conn, :atex_session) + + case oauth_did do + did when is_binary(did) -> + client = %__MODULE__{did: did} + + with_session_lock(client, fn -> + case maybe_refresh(client) do + {:ok, _session} -> {:ok, client} + _ -> {:error, :reauth} + end + end) _ -> :error end end - @doc """ - Updates a `Plug.Conn` session with the latest values from the client. - - Ideally should be called at the end of routes where XRPC calls occur, in case - the client has transparently refreshed, so that the user is always up to date. - """ - @spec update_conn(Plug.Conn.t(), t()) :: Plug.Conn.t() - def update_conn(%Plug.Conn{} = conn, %__MODULE__{} = client) do - Plug.Conn.put_session(conn, :atex_oauth, %{ - access_token: client.access_token, - refresh_token: client.refresh_token, - did: client.did, - pds: client.endpoint, - expires_at: client.expires_at, - dpop_nonce: client.dpop_nonce, - dpop_key: client.dpop_key - }) - end - @doc """ Ask the client's OAuth server for a new set of auth tokens. + Fetches the session, refreshes the tokens, creates a new session with the + updated tokens, stores it, and returns the new session. + You shouldn't need to call this manually for the most part, the client does - it's best to refresh automatically when it needs to. + its best to refresh automatically when it needs to. + + This function acquires a lock on the session to prevent concurrent refresh attempts. """ - @spec refresh(t()) :: {:ok, t()} | {:error, any()} + @spec refresh(client :: t()) :: {:ok, OAuth.Session.t()} | {:error, any()} def refresh(%__MODULE__{} = client) do - with {:ok, authz_server} <- OAuth.get_authorization_server(client.endpoint), + with_session_lock(client, fn -> + do_refresh(client) + end) + end + + @spec do_refresh(t()) :: {:ok, OAuth.Session.t()} | {:error, any()} + defp do_refresh(%__MODULE__{did: did}) do + with {:ok, session} <- OAuth.SessionStore.get(did), + {:ok, authz_server} <- OAuth.get_authorization_server(session.aud), {:ok, %{token_endpoint: token_endpoint}} <- OAuth.get_authorization_server_metadata(authz_server) do case OAuth.refresh_token( - client.refresh_token, - client.dpop_key, - client.issuer, + session.refresh_token, + session.dpop_key, + session.iss, token_endpoint ) do {:ok, tokens, nonce} -> - {:ok, - %{ - client - | access_token: tokens.access_token, - refresh_token: tokens.refresh_token, - dpop_nonce: nonce - }} + new_session = %OAuth.Session{ + iss: session.iss, + aud: session.aud, + sub: tokens.did, + access_token: tokens.access_token, + refresh_token: tokens.refresh_token, + expires_at: tokens.expires_at, + dpop_key: session.dpop_key, + dpop_nonce: nonce + } + + case OAuth.SessionStore.update(new_session) do + :ok -> {:ok, new_session} + err -> err + end err -> err @@ -124,72 +153,119 @@ defmodule Atex.XRPC.OAuthClient do end end + @spec maybe_refresh(t(), integer()) :: {:ok, OAuth.Session.t()} | {:error, any()} + defp maybe_refresh(%__MODULE__{did: did} = client, buffer_minutes \\ 5) do + with {:ok, session} <- OAuth.SessionStore.get(did) do + if token_expiring_soon?(session.expires_at, buffer_minutes) do + do_refresh(client) + else + {:ok, session} + end + end + end + + @spec token_expiring_soon?(NaiveDateTime.t(), integer()) :: boolean() + defp token_expiring_soon?(expires_at, buffer_minutes) do + now = NaiveDateTime.utc_now() + expiry_threshold = NaiveDateTime.add(now, buffer_minutes * 60, :second) + + NaiveDateTime.compare(expires_at, expiry_threshold) in [:lt, :eq] + end + @doc """ - See `Atex.XRPC.get/3`. + Make a GET request to an XRPC endpoint. + + See `Atex.XRPC.get/3` for details. """ @impl true def get(%__MODULE__{} = client, resource, opts \\ []) do - request(client, opts ++ [method: :get, url: XRPC.url(client.endpoint, resource)]) + # TODO: Keyword.valiate to make sure :method isn't passed? + request(client, resource, opts ++ [method: :get]) end @doc """ - See `Atex.XRPC.post/3`. + Make a POST request to an XRPC endpoint. + + See `Atex.XRPC.post/3` for details. """ @impl true def post(%__MODULE__{} = client, resource, opts \\ []) do - request(client, opts ++ [method: :post, url: XRPC.url(client.endpoint, resource)]) + # Ditto + request(client, resource, opts ++ [method: :post]) end - @spec request(t(), keyword()) :: {:ok, Req.Response.t(), t()} | {:error, any(), any()} - defp request(client, opts) do - # Preemptively refresh token if it's about to expire - with {:ok, client} <- maybe_refresh(client) do - request = opts |> Req.new() |> put_auth(client.access_token) - - case OAuth.request_protected_dpop_resource( - request, - client.issuer, - client.access_token, - client.dpop_key, - client.dpop_nonce - ) do - {:ok, %{status: 200} = response, nonce} -> - client = %{client | dpop_nonce: nonce} - {:ok, response, client} + defp request(%__MODULE__{} = client, resource, opts) do + with_session_lock(client, fn -> + case maybe_refresh(client) do + {:ok, session} -> + url = Atex.XRPC.url(session.aud, resource) + + request = + opts + |> Keyword.put(:url, url) + |> Req.new() + |> Req.Request.put_header("authorization", "DPoP #{session.access_token}") - {:ok, response, nonce} -> - client = %{client | dpop_nonce: nonce} - handle_failure(client, response, request) + case OAuth.request_protected_dpop_resource( + request, + session.iss, + session.access_token, + session.dpop_key, + session.dpop_nonce + ) do + {:ok, %{status: 200} = response, nonce} -> + update_session_nonce(session, nonce) + {:ok, response, client} + + {:ok, response, nonce} -> + update_session_nonce(session, nonce) + handle_failure(client, request, response) + + err -> + err + end err -> err end - end + end) end - @spec handle_failure(t(), Req.Response.t(), Req.Request.t()) :: - {:ok, Req.Response.t(), t()} | {:error, any(), t()} - defp handle_failure(client, response, request) do - IO.inspect(response, label: "got failure") + # Execute a function with an exclusive lock on the session identified by the + # client's DID. This ensures that concurrent requests for the same user don't + # race during token refresh. + @spec with_session_lock(t(), (-> result)) :: result when result: any() + defp with_session_lock(%__MODULE__{did: did}, fun) do + Mutex.with_lock(Atex.SessionMutex, did, fun) + end - if auth_error?(response.body) and client.refresh_token do - case refresh(client) do - {:ok, client} -> + defp handle_failure(client, request, response) do + if auth_error?(response) do + case do_refresh(client) do + {:ok, session} -> case OAuth.request_protected_dpop_resource( request, - client.issuer, - client.access_token, - client.dpop_key, - client.dpop_nonce + session.iss, + session.access_token, + session.dpop_key, + session.dpop_nonce ) do {:ok, %{status: 200} = response, nonce} -> - {:ok, response, %{client | dpop_nonce: nonce}} + update_session_nonce(session, nonce) + {:ok, response, client} - {:ok, response, nonce} -> - {:error, response, %{client | dpop_nonce: nonce}} + {:ok, response, _nonce} -> + if auth_error?(response) do + # We tried to refresh the token once but it's still failing + # Clear session and prompt dev to reauth or something + OAuth.SessionStore.delete(session) + {:error, response, :expired} + else + {:error, response, client} + end - {:error, err} -> - {:error, err, client} + err -> + err end err -> @@ -200,29 +276,17 @@ defmodule Atex.XRPC.OAuthClient do end end - @spec maybe_refresh(t(), integer()) :: {:ok, t()} | {:error, any()} - defp maybe_refresh(%__MODULE__{expires_at: expires_at} = client, buffer_minutes \\ 5) do - if token_expiring_soon?(expires_at, buffer_minutes) do - refresh(client) - else - {:ok, client} - end - end + @spec auth_error?(Req.Response.t()) :: boolean() + defp auth_error?(%{status: 401, headers: %{"www-authenticate" => [www_auth]}}), + do: + (String.starts_with?(www_auth, "Bearer") or String.starts_with?(www_auth, "DPoP")) and + String.contains?(www_auth, "error=\"invalid_token\"") - @spec token_expiring_soon?(NaiveDateTime.t(), integer()) :: boolean() - defp token_expiring_soon?(expires_at, buffer_minutes) do - now = NaiveDateTime.utc_now() - expiry_threshold = NaiveDateTime.add(now, buffer_minutes * 60, :second) + defp auth_error?(_resp), do: false - NaiveDateTime.compare(expires_at, expiry_threshold) in [:lt, :eq] + defp update_session_nonce(session, nonce) do + session = %{session | dpop_nonce: nonce} + :ok = OAuth.SessionStore.update(session) + session end - - @spec auth_error?(body :: Req.Response.t()) :: boolean() - defp auth_error?(%{status: status}) when status in [401, 403], do: true - defp auth_error?(%{body: %{"error" => "InvalidToken"}}), do: true - defp auth_error?(_response), do: false - - @spec put_auth(Req.Request.t(), String.t()) :: Req.Request.t() - defp put_auth(request, token), - do: Req.Request.put_header(request, "authorization", "DPoP #{token}") end diff --git a/mix.exs b/mix.exs index bcd9c13..b0e155e 100644 --- a/mix.exs +++ b/mix.exs @@ -41,7 +41,8 @@ defmodule Atex.MixProject do {:jose, "~> 1.11"}, {:bandit, "~> 1.0", only: [:dev, :test]}, {:con_cache, "~> 1.1"}, - {:mutex, "~> 3.0"} + {:mutex, "~> 3.0"}, + {:dialyxir, "~> 1.4", only: [:dev, :test], runtime: false} ] end diff --git a/mix.lock b/mix.lock index 9672e9e..41485eb 100644 --- a/mix.lock +++ b/mix.lock @@ -5,7 +5,9 @@ "con_cache": {:hex, :con_cache, "1.1.1", "9f47a68dfef5ac3bbff8ce2c499869dbc5ba889dadde6ac4aff8eb78ddaf6d82", [:mix], [{:telemetry, "~> 1.0", [hex: :telemetry, repo: "hexpm", optional: false]}], "hexpm", "1def4d1bec296564c75b5bbc60a19f2b5649d81bfa345a2febcc6ae380e8ae15"}, "credo": {:hex, :credo, "1.7.15", "283da72eeb2fd3ccf7248f4941a0527efb97afa224bcdef30b4b580bc8258e1c", [:mix], [{:bunt, "~> 0.2.1 or ~> 1.0", [hex: :bunt, repo: "hexpm", optional: false]}, {:file_system, "~> 0.2 or ~> 1.0", [hex: :file_system, repo: "hexpm", optional: false]}, {:jason, "~> 1.0", [hex: :jason, repo: "hexpm", optional: false]}], "hexpm", "291e8645ea3fea7481829f1e1eb0881b8395db212821338e577a90bf225c5607"}, "decimal": {:hex, :decimal, "2.3.0", "3ad6255aa77b4a3c4f818171b12d237500e63525c2fd056699967a3e7ea20f62", [:mix], [], "hexpm", "a4d66355cb29cb47c3cf30e71329e58361cfcb37c34235ef3bf1d7bf3773aeac"}, + "dialyxir": {:hex, :dialyxir, "1.4.7", "dda948fcee52962e4b6c5b4b16b2d8fa7d50d8645bbae8b8685c3f9ecb7f5f4d", [:mix], [{:erlex, ">= 0.2.8", [hex: :erlex, repo: "hexpm", optional: false]}], "hexpm", "b34527202e6eb8cee198efec110996c25c5898f43a4094df157f8d28f27d9efe"}, "earmark_parser": {:hex, :earmark_parser, "1.4.44", "f20830dd6b5c77afe2b063777ddbbff09f9759396500cdbe7523efd58d7a339c", [:mix], [], "hexpm", "4778ac752b4701a5599215f7030989c989ffdc4f6df457c5f36938cc2d2a2750"}, + "erlex": {:hex, :erlex, "0.2.8", "cd8116f20f3c0afe376d1e8d1f0ae2452337729f68be016ea544a72f767d9c12", [:mix], [], "hexpm", "9d66ff9fedf69e49dc3fd12831e12a8a37b76f8651dd21cd45fcf5561a8a7590"}, "ex_cldr": {:hex, :ex_cldr, "2.44.1", "0d220b175874e1ce77a0f7213bdfe700b9be11aefbf35933a0e98837803ebdc5", [:mix], [{:cldr_utils, "~> 2.28", [hex: :cldr_utils, repo: "hexpm", optional: false]}, {:decimal, "~> 1.6 or ~> 2.0", [hex: :decimal, repo: "hexpm", optional: false]}, {:gettext, "~> 0.19 or ~> 1.0", [hex: :gettext, repo: "hexpm", optional: true]}, {:jason, "~> 1.0", [hex: :jason, repo: "hexpm", optional: true]}, {:nimble_parsec, "~> 0.5 or ~> 1.0", [hex: :nimble_parsec, repo: "hexpm", optional: true]}], "hexpm", "3880cd6137ea21c74250cd870d3330c4a9fdec07fabd5e37d1b239547929e29b"}, "ex_doc": {:hex, :ex_doc, "0.39.3", "519c6bc7e84a2918b737aec7ef48b96aa4698342927d080437f61395d361dcee", [:mix], [{:earmark_parser, "~> 1.4.44", [hex: :earmark_parser, repo: "hexpm", optional: false]}, {:makeup_c, ">= 0.1.0", [hex: :makeup_c, repo: "hexpm", optional: true]}, {:makeup_elixir, "~> 0.14 or ~> 1.0", [hex: :makeup_elixir, repo: "hexpm", optional: false]}, {:makeup_erlang, "~> 0.1 or ~> 1.0", [hex: :makeup_erlang, repo: "hexpm", optional: false]}, {:makeup_html, ">= 0.1.0", [hex: :makeup_html, repo: "hexpm", optional: true]}], "hexpm", "0590955cf7ad3b625780ee1c1ea627c28a78948c6c0a9b0322bd976a079996e1"}, "file_system": {:hex, :file_system, "1.1.1", "31864f4685b0148f25bd3fbef2b1228457c0c89024ad67f7a81a3ffbc0bbad3a", [:mix], [], "hexpm", "7a15ff97dfe526aeefb090a7a9d3d03aa907e100e262a0f8f7746b78f8f87a5d"},