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 mut session_service = systemprompt_oauth::services::SessionCreationService::new(
185        Arc::clone(state.analytics_provider()),
186        Arc::clone(state.user_provider()),
187    );
188    if let Some(publisher) = state.event_publisher() {
189        session_service = session_service.with_event_publisher(Arc::clone(publisher));
190    }
191    let analytics = state.analytics_provider().extract_analytics(
192        params.headers,
193        ExtractSignals {
194            caller_ip: params.caller_ip,
195            ..Default::default()
196        },
197    );
198    let session_id = session_service
199        .create_authenticated_session(&validation_result.user_id, &analytics, SessionSource::Oauth)
200        .await?;
201
202    let access_token_jti = generate_access_token_jti();
203    let global_config = Config::get()?;
204    let config = JwtConfig {
205        permissions: permissions.clone(),
206        audience: global_config.jwt_audiences.clone(),
207        ..Default::default()
208    };
209    let signing = JwtSigningParams {
210        issuer: &global_config.jwt_issuer,
211    };
212    let access_token = generate_jwt(&user, config, access_token_jti, &session_id, &signing)?;
213
214    let refresh_token_value = generate_secure_token("rt");
215    let refresh_token_id = RefreshTokenId::new(&refresh_token_value);
216    let refresh_expires_at = chrono::Utc::now().timestamp()
217        + (systemprompt_oauth::constants::token::SECONDS_PER_DAY
218            * systemprompt_oauth::constants::token::REFRESH_TOKEN_EXPIRY_DAYS);
219
220    let refresh_params = RefreshTokenParams::builder(
221        &refresh_token_id,
222        params.client_id,
223        &validation_result.user_id,
224        &validation_result.scope,
225        refresh_expires_at,
226    )
227    .build();
228    repo.store_refresh_token(refresh_params).await?;
229
230    if let Err(e) = repo
231        .link_auth_code_to_refresh_token(params.code, refresh_token_id.as_str())
232        .await
233    {
234        tracing::warn!(error = %e, "Failed to link auth code to refresh token");
235    }
236
237    Ok(TokenResponse { access_token })
238}
239
240async fn load_authenticated_user(
241    user_id: &UserId,
242    user_provider: &Arc<dyn systemprompt_traits::UserProvider>,
243) -> anyhow::Result<AuthenticatedUser> {
244    let user = user_provider
245        .find_by_id(user_id)
246        .await
247        .map_err(|e| anyhow::anyhow!("{}", e))?
248        .ok_or_else(|| anyhow::anyhow!("User not found: {user_id}"))?;
249
250    let permissions: Vec<Permission> = user
251        .roles
252        .iter()
253        .filter_map(|s| {
254            Permission::from_str(s)
255                .map_err(|e| {
256                    tracing::warn!(
257                        user_id = %user.id,
258                        role = %s,
259                        error = %e,
260                        "Invalid role in user record"
261                    );
262                    e
263                })
264                .ok()
265        })
266        .collect();
267
268    let user_uuid = uuid::Uuid::parse_str(user.id.as_str())
269        .map_err(|_e| anyhow::anyhow!("Invalid user UUID: {}", user.id))?;
270
271    Ok(AuthenticatedUser::new_with_roles(
272        user_uuid,
273        user.name,
274        user.email,
275        permissions,
276        user.roles,
277    ))
278}
279
280#[derive(Debug)]
281struct BrowserClient {
282    client_id: String,
283}
284
285#[derive(Debug, serde::Deserialize)]
286struct TokenResponse {
287    access_token: String,
288}