Skip to main content

systemprompt_api/routes/oauth/endpoints/
callback.rs

1//! OAuth callback endpoint for the server's own browser client.
2//!
3//! Exchanges the returned authorization code for tokens, establishes an
4//! authenticated session, sets the access-token cookie, and redirects to the
5//! origin-validated `return_to` recovered from the consumed state binding.
6//!
7//! Copyright (c) systemprompt.io — Business Source License 1.1.
8//! See <https://systemprompt.io> for licensing details.
9
10use axum::extract::{Query, State};
11use axum::http::{HeaderMap, HeaderValue, StatusCode, header};
12use axum::response::{IntoResponse, Redirect, Response};
13use serde::Deserialize;
14use std::str::FromStr;
15use std::sync::Arc;
16use systemprompt_identifiers::{
17    AuthorizationCode, ClientId, RefreshTokenId, SessionSource, UserId,
18};
19use systemprompt_models::Config;
20use systemprompt_models::auth::{AuthenticatedUser, Permission, parse_permissions};
21
22use crate::routes::oauth::extractors::OAuthRepo;
23use crate::services::middleware::client_addr::ClientIp;
24use systemprompt_oauth::OAuthState;
25use systemprompt_oauth::repository::{OAuthRepository, RefreshTokenParams};
26use systemprompt_traits::ExtractSignals;
27
28#[derive(Debug, Deserialize)]
29pub struct CallbackQuery {
30    pub code: String,
31    pub state: Option<String>,
32}
33
34pub async fn handle_callback(
35    Query(params): Query<CallbackQuery>,
36    State(state): State<OAuthState>,
37    OAuthRepo(repo): OAuthRepo,
38    ClientIp(caller_ip): ClientIp,
39    headers: HeaderMap,
40) -> impl IntoResponse {
41    let config = match Config::get() {
42        Ok(c) => c,
43        Err(e) => {
44            return (
45                StatusCode::INTERNAL_SERVER_ERROR,
46                format!("Failed to load config: {e}"),
47            )
48                .into_response();
49        },
50    };
51
52    let server_base_url = &config.api_external_url;
53    let redirect_uri = format!("{server_base_url}/api/v1/core/oauth/callback");
54
55    let browser_client = match find_browser_client(&repo, &redirect_uri).await {
56        Ok(client) => client,
57        Err(e) => {
58            return (
59                StatusCode::INTERNAL_SERVER_ERROR,
60                format!("Failed to find OAuth client: {e}"),
61            )
62                .into_response();
63        },
64    };
65
66    let code = AuthorizationCode::new(&params.code);
67    let client_id = ClientId::new(&browser_client.client_id);
68    let token_response = match exchange_code_for_token(
69        &repo,
70        CodeExchangeParams {
71            caller_ip,
72            code: &code,
73            client_id: &client_id,
74            redirect_uri: &redirect_uri,
75            headers: &headers,
76        },
77        &state,
78    )
79    .await
80    {
81        Ok(response) => response,
82        Err(e) => {
83            return (
84                StatusCode::UNAUTHORIZED,
85                format!("Failed to exchange code for token: {e}"),
86            )
87                .into_response();
88        },
89    };
90
91    let redirect_destination =
92        match resolve_redirect_destination(&repo, params.state.as_deref()).await {
93            Ok(destination) => destination,
94            Err(response) => return response,
95        };
96
97    session_cookie_redirect(&token_response.access_token, &redirect_destination)
98}
99
100async fn resolve_redirect_destination(
101    repo: &OAuthRepository,
102    state_token: Option<&str>,
103) -> Result<String, Response> {
104    let Some(state_token) = state_token.filter(|s| !s.is_empty()) else {
105        return Err((StatusCode::BAD_REQUEST, "Missing state parameter").into_response());
106    };
107    match repo.consume_state_binding(state_token).await {
108        Ok(Some(binding)) => Ok(binding.return_to),
109        Ok(None) => {
110            tracing::warn!("state binding missing, expired, or already consumed");
111            Err((StatusCode::BAD_REQUEST, "Invalid state parameter").into_response())
112        },
113        Err(e) => {
114            tracing::error!(error = %e, "state binding lookup failed");
115            Err((
116                StatusCode::INTERNAL_SERVER_ERROR,
117                "Failed to validate state",
118            )
119                .into_response())
120        },
121    }
122}
123
124fn session_cookie_redirect(access_token: &str, destination: &str) -> Response {
125    let cookie = format!(
126        "access_token={access_token}; Path=/; HttpOnly; Secure; SameSite=Strict; Max-Age={}",
127        systemprompt_oauth::constants::token::COOKIE_MAX_AGE_SECONDS
128    );
129
130    let mut response = Redirect::to(destination).into_response();
131    if let Ok(cookie_value) = HeaderValue::from_str(&cookie) {
132        response
133            .headers_mut()
134            .insert(header::SET_COOKIE, cookie_value);
135    }
136
137    response
138}
139
140async fn find_browser_client(
141    repo: &OAuthRepository,
142    redirect_uri: &str,
143) -> anyhow::Result<BrowserClient> {
144    let client = repo
145        .find_client_by_redirect_uri_with_scope(redirect_uri, &["admin", "user"])
146        .await?
147        .ok_or_else(|| anyhow::anyhow!("No suitable browser client found"))?;
148
149    Ok(BrowserClient {
150        client_id: client.client_id.to_string(),
151    })
152}
153
154struct CodeExchangeParams<'a> {
155    caller_ip: Option<std::net::IpAddr>,
156    code: &'a AuthorizationCode,
157    client_id: &'a ClientId,
158    redirect_uri: &'a str,
159    headers: &'a HeaderMap,
160}
161
162async fn exchange_code_for_token(
163    repo: &OAuthRepository,
164    params: CodeExchangeParams<'_>,
165    state: &OAuthState,
166) -> anyhow::Result<TokenResponse> {
167    use systemprompt_oauth::services::{
168        JwtConfig, JwtSigningParams, generate_access_token_jti, generate_jwt, generate_secure_token,
169    };
170
171    let validation_result = repo
172        .validate_authorization_code(
173            params.code,
174            params.client_id,
175            Some(params.redirect_uri),
176            None,
177        )
178        .await?;
179
180    let user = load_authenticated_user(&validation_result.user_id, state.user_provider()).await?;
181
182    let permissions = parse_permissions(&validation_result.scope)?;
183
184    let session_service = systemprompt_oauth::services::SessionCreationService::new(
185        Arc::clone(state.session_provider()),
186        Arc::clone(state.user_provider()),
187    );
188    let analytics = state.analytics_provider().extract_analytics(
189        params.headers,
190        ExtractSignals {
191            caller_ip: params.caller_ip,
192            ..Default::default()
193        },
194    );
195    let session_id = session_service
196        .create_authenticated_session(&validation_result.user_id, &analytics, SessionSource::Oauth)
197        .await?;
198
199    let access_token_jti = generate_access_token_jti();
200    let global_config = Config::get()?;
201    let config = JwtConfig {
202        permissions: permissions.clone(),
203        audience: global_config.jwt_audiences.clone(),
204        client_id: Some(params.client_id.clone()),
205        ..Default::default()
206    };
207    let signing = JwtSigningParams {
208        issuer: &global_config.jwt_issuer,
209    };
210    let access_token = generate_jwt(&user, config, access_token_jti, &session_id, &signing)?;
211
212    let refresh_token_value = generate_secure_token("rt");
213    let refresh_token_id = RefreshTokenId::new(&refresh_token_value);
214    let refresh_expires_at = chrono::Utc::now().timestamp()
215        + (systemprompt_oauth::constants::token::SECONDS_PER_DAY
216            * systemprompt_oauth::constants::token::REFRESH_TOKEN_EXPIRY_DAYS);
217
218    let refresh_params = RefreshTokenParams::builder(
219        &refresh_token_id,
220        params.client_id,
221        &validation_result.user_id,
222        &validation_result.scope,
223        refresh_expires_at,
224    )
225    .build();
226    repo.store_refresh_token(refresh_params).await?;
227
228    if let Err(e) = repo
229        .link_auth_code_to_refresh_token(params.code, refresh_token_id.as_str())
230        .await
231    {
232        tracing::warn!(error = %e, "Failed to link auth code to refresh token");
233    }
234
235    Ok(TokenResponse { access_token })
236}
237
238async fn load_authenticated_user(
239    user_id: &UserId,
240    user_provider: &Arc<dyn systemprompt_traits::UserProvider>,
241) -> anyhow::Result<AuthenticatedUser> {
242    let user = user_provider
243        .find_by_id(user_id)
244        .await
245        .map_err(|e| anyhow::anyhow!("{}", e))?
246        .ok_or_else(|| anyhow::anyhow!("User not found: {user_id}"))?;
247
248    let permissions: Vec<Permission> = user
249        .roles
250        .iter()
251        .filter_map(|s| {
252            Permission::from_str(s)
253                .map_err(|e| {
254                    tracing::warn!(
255                        user_id = %user.id,
256                        role = %s,
257                        error = %e,
258                        "Invalid role in user record"
259                    );
260                    e
261                })
262                .ok()
263        })
264        .collect();
265
266    let user_uuid = uuid::Uuid::parse_str(user.id.as_str())
267        .map_err(|_e| anyhow::anyhow!("Invalid user UUID: {}", user.id))?;
268
269    Ok(AuthenticatedUser::new_with_roles(
270        user_uuid,
271        user.name,
272        user.email,
273        permissions,
274        user.roles,
275    ))
276}
277
278#[derive(Debug)]
279struct BrowserClient {
280    client_id: String,
281}
282
283#[derive(Debug, serde::Deserialize)]
284struct TokenResponse {
285    access_token: String,
286}