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::sync::Arc;
15use systemprompt_identifiers::{
16    AccessTokenId, AuthorizationCode, ClientId, RefreshTokenId, SessionSource,
17};
18use systemprompt_manifest::Config;
19use systemprompt_models::auth::parse_permissions;
20
21use crate::routes::oauth::extractors::OAuthRepo;
22use crate::routes::oauth::{OAuthHttpError, internal};
23use crate::services::middleware::client_addr::ClientIp;
24use systemprompt_oauth::OAuthState;
25use systemprompt_oauth::repository::{OAuthRepository, RefreshTokenParams};
26use systemprompt_oauth::services::load_authenticated_user;
27use systemprompt_traits::ExtractSignals;
28
29#[derive(Debug, Deserialize)]
30pub struct CallbackQuery {
31    pub code: String,
32    pub state: Option<String>,
33}
34
35pub async fn handle_callback(
36    Query(params): Query<CallbackQuery>,
37    State(state): State<OAuthState>,
38    OAuthRepo(repo): OAuthRepo,
39    ClientIp(caller_ip): ClientIp,
40    headers: HeaderMap,
41) -> Result<Response, OAuthHttpError> {
42    let config = Config::get()?;
43
44    let server_base_url = &config.api_external_url;
45    let redirect_uri = format!("{server_base_url}/api/v1/core/oauth/callback");
46
47    let browser_client = find_browser_client(&repo, &redirect_uri)
48        .await
49        .map_err(|e| internal::server_error("Failed to find OAuth client", e))?;
50
51    let code = AuthorizationCode::new(&params.code);
52    let client_id = ClientId::new(&browser_client.client_id);
53    let token_response = exchange_code_for_token(
54        &repo,
55        CodeExchangeParams {
56            caller_ip,
57            code: &code,
58            client_id: &client_id,
59            redirect_uri: &redirect_uri,
60            headers: &headers,
61        },
62        &state,
63    )
64    .await
65    .map_err(|e| {
66        internal::rejected(
67            OAuthHttpError::invalid_grant("Failed to exchange code for token")
68                .with_status(StatusCode::UNAUTHORIZED),
69            e,
70        )
71    })?;
72
73    let redirect_destination = resolve_redirect_destination(&repo, params.state.as_deref()).await?;
74
75    Ok(session_cookie_redirect(
76        &token_response.access_token,
77        &redirect_destination,
78    ))
79}
80
81async fn resolve_redirect_destination(
82    repo: &OAuthRepository,
83    state_token: Option<&str>,
84) -> Result<String, OAuthHttpError> {
85    let Some(state_token) = state_token.filter(|s| !s.is_empty()) else {
86        return Err(OAuthHttpError::invalid_request("Missing state parameter"));
87    };
88    match repo.consume_state_binding(state_token).await {
89        Ok(Some(binding)) => Ok(binding.return_to),
90        Ok(None) => {
91            tracing::warn!("state binding missing, expired, or already consumed");
92            Err(OAuthHttpError::invalid_request("Invalid state parameter"))
93        },
94        Err(e) => Err(internal::server_error("Failed to validate state", e)),
95    }
96}
97
98fn session_cookie_redirect(access_token: &str, destination: &str) -> Response {
99    let cookie = format!(
100        "access_token={access_token}; Path=/; HttpOnly; Secure; SameSite=Strict; Max-Age={}",
101        systemprompt_oauth::constants::token::COOKIE_MAX_AGE_SECONDS
102    );
103
104    let mut response = Redirect::to(destination).into_response();
105    if let Ok(cookie_value) = HeaderValue::from_str(&cookie) {
106        response
107            .headers_mut()
108            .insert(header::SET_COOKIE, cookie_value);
109    }
110
111    response
112}
113
114async fn find_browser_client(
115    repo: &OAuthRepository,
116    redirect_uri: &str,
117) -> anyhow::Result<BrowserClient> {
118    let client = repo
119        .find_client_by_redirect_uri_with_scope(redirect_uri, &["admin", "user"])
120        .await?
121        .ok_or_else(|| anyhow::anyhow!("No suitable browser client found"))?;
122
123    Ok(BrowserClient {
124        client_id: client.client_id.to_string(),
125    })
126}
127
128struct CodeExchangeParams<'a> {
129    caller_ip: Option<std::net::IpAddr>,
130    code: &'a AuthorizationCode,
131    client_id: &'a ClientId,
132    redirect_uri: &'a str,
133    headers: &'a HeaderMap,
134}
135
136async fn exchange_code_for_token(
137    repo: &OAuthRepository,
138    params: CodeExchangeParams<'_>,
139    state: &OAuthState,
140) -> anyhow::Result<TokenResponse> {
141    use systemprompt_oauth::services::{
142        JwtConfig, JwtSigningParams, generate_jwt, generate_secure_token,
143    };
144
145    let validation_result = repo
146        .validate_authorization_code(params.code, params.client_id, params.redirect_uri, "")
147        .await?;
148
149    let user =
150        load_authenticated_user(state.user_provider().as_ref(), &validation_result.user_id).await?;
151
152    let permissions = parse_permissions(&validation_result.scope)?;
153
154    let session_service = systemprompt_oauth::services::SessionCreationService::new(
155        Arc::clone(state.session_provider()),
156        Arc::clone(state.user_provider()),
157    );
158    let analytics = state.analytics_provider().extract_analytics(
159        params.headers,
160        ExtractSignals {
161            caller_ip: params.caller_ip,
162            ..Default::default()
163        },
164    );
165    let session_id = session_service
166        .create_authenticated_session(&validation_result.user_id, &analytics, SessionSource::Oauth)
167        .await?;
168
169    let access_token_jti = AccessTokenId::generate();
170    let global_config = Config::get()?;
171    let config = JwtConfig {
172        permissions: permissions.clone(),
173        audience: global_config.jwt_audiences.clone(),
174        client_id: Some(params.client_id.clone()),
175        ..Default::default()
176    };
177    let signing = JwtSigningParams {
178        issuer: &global_config.jwt_issuer,
179    };
180    let access_token = generate_jwt(&user, config, access_token_jti, &session_id, &signing)?;
181
182    let refresh_token_value = generate_secure_token("rt");
183    let refresh_token_id = RefreshTokenId::new(&refresh_token_value);
184    let refresh_expires_at = chrono::Utc::now().timestamp()
185        + (systemprompt_oauth::constants::token::SECONDS_PER_DAY
186            * systemprompt_oauth::constants::token::REFRESH_TOKEN_EXPIRY_DAYS);
187
188    let refresh_params = RefreshTokenParams::builder(
189        &refresh_token_id,
190        params.client_id,
191        &validation_result.user_id,
192        &validation_result.scope,
193        refresh_expires_at,
194    )
195    .build();
196    repo.store_refresh_token(refresh_params).await?;
197
198    if let Err(e) = repo
199        .link_auth_code_to_refresh_token(params.code, &refresh_token_id)
200        .await
201    {
202        tracing::warn!(error = %e, "Failed to link auth code to refresh token");
203    }
204
205    Ok(TokenResponse { access_token })
206}
207
208#[derive(Debug)]
209struct BrowserClient {
210    client_id: String,
211}
212
213#[derive(Debug, serde::Deserialize)]
214struct TokenResponse {
215    access_token: String,
216}