systemprompt_api/routes/oauth/endpoints/
callback.rs1use 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(¶ms.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}