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 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}