backbone_integrations/presentation/http/
oauth_handler.rs1use std::sync::Arc;
36
37use axum::extract::{Path, Query, State};
38use axum::http::{header, HeaderMap, StatusCode};
39use axum::middleware::{self, Next};
40use axum::response::{Html, IntoResponse, Response};
41use axum::routing::{get, post};
42use axum::{Json, Router};
43use serde::Deserialize;
44use uuid::Uuid;
45
46use crate::application::service::integrations_oauth::{
47 AuthorizeRequest, AuthorizeResponse, CompleteOutcome, CompleteRequest, IntegrationsOauthService,
48 OauthError,
49};
50
51#[derive(Debug, Clone)]
60pub struct OAuthPrincipal {
61 pub company_id: Uuid,
62 pub user_id: Option<Uuid>,
63 pub permissions: Vec<String>,
64}
65
66impl OAuthPrincipal {
67 fn has(&self, permission: &str) -> bool {
68 self.permissions.iter().any(|p| p == permission)
69 }
70}
71
72pub async fn require_principal(req: axum::extract::Request, next: Next) -> Response {
75 if req.extensions().get::<OAuthPrincipal>().is_none() {
76 return error_response(
77 StatusCode::UNAUTHORIZED,
78 "OAUTH_UNAUTHENTICATED",
79 "authentication required: no validated principal on the request",
80 );
81 }
82 next.run(req).await
83}
84
85fn error_response(status: StatusCode, code: &str, message: impl std::fmt::Display) -> Response {
90 (
91 status,
92 Json(serde_json::json!({
93 "success": false,
94 "error": code,
95 "message": message.to_string(),
96 })),
97 )
98 .into_response()
99}
100
101impl IntoResponse for OauthError {
102 fn into_response(self) -> Response {
103 let (status, code) = match &self {
104 OauthError::Invalid(_) => (StatusCode::BAD_REQUEST, "OAUTH_INVALID_INPUT"),
105 OauthError::State(_) => (StatusCode::BAD_REQUEST, "OAUTH_STATE_REJECTED"),
106 OauthError::Identity(_) => (StatusCode::BAD_REQUEST, "OAUTH_IDENTITY_REJECTED"),
107 OauthError::Unstoreable(_) => (StatusCode::BAD_GATEWAY, "OAUTH_UNSTOREABLE_TOKEN"),
108 OauthError::NotFound => (StatusCode::NOT_FOUND, "OAUTH_ACCOUNT_NOT_FOUND"),
109 OauthError::ProviderUnconfigured(_) => {
110 (StatusCode::SERVICE_UNAVAILABLE, "OAUTH_PROVIDER_UNCONFIGURED")
111 }
112 OauthError::Transport(_) => (StatusCode::BAD_GATEWAY, "OAUTH_PROVIDER_TRANSPORT"),
113 OauthError::Store(_) => (StatusCode::SERVICE_UNAVAILABLE, "OAUTH_CREDENTIAL_STORE"),
114 OauthError::Db(_) => (StatusCode::INTERNAL_SERVER_ERROR, "OAUTH_DATABASE"),
115 };
116 error_response(status, code, self)
117 }
118}
119
120fn forbidden(permission: &str) -> Response {
121 error_response(
122 StatusCode::FORBIDDEN,
123 "OAUTH_FORBIDDEN",
124 format!("this route requires the {permission:?} permission"),
125 )
126}
127
128async fn authorize(
134 State(service): State<Arc<IntegrationsOauthService>>,
135 axum::Extension(principal): axum::Extension<OAuthPrincipal>,
136 Json(req): Json<AuthorizeRequest>,
137) -> Response {
138 if !principal.has("write:integrations") {
139 return forbidden("write:integrations");
140 }
141 match service.authorize(req).await {
145 Ok(out) => (StatusCode::OK, Json(serde_json::json!({ "success": true, "data": out }))).into_response(),
146 Err(e) => e.into_response(),
147 }
148}
149
150#[derive(Debug, Deserialize)]
153struct CallbackQuery {
154 code: Option<String>,
155 state: Option<String>,
156 error: Option<String>,
157}
158
159async fn callback(
163 State(service): State<Arc<IntegrationsOauthService>>,
164 Query(q): Query<CallbackQuery>,
165) -> Response {
166 if let Some(err) = q.error.as_deref().filter(|e| !e.trim().is_empty()) {
167 return error_page(StatusCode::BAD_REQUEST, "the provider refused the authorization", err);
168 }
169 let (Some(code), Some(state)) = (q.code.as_deref(), q.state.as_deref()) else {
170 return error_page(
171 StatusCode::BAD_REQUEST,
172 "incomplete callback",
173 "the callback is missing its code or signed state",
174 );
175 };
176 match service.callback_page(code, state) {
177 Ok(page) => Html(page).into_response(),
178 Err(e) => error_page(StatusCode::BAD_REQUEST, "the connection could not continue", &e.to_string()),
179 }
180}
181
182async fn complete(
186 State(service): State<Arc<IntegrationsOauthService>>,
187 axum::Extension(principal): axum::Extension<OAuthPrincipal>,
188 headers: HeaderMap,
189 body: axum::body::Bytes,
190) -> Response {
191 if !principal.has("write:integrations") {
192 return forbidden("write:integrations");
193 }
194 let content_type = headers
195 .get(header::CONTENT_TYPE)
196 .and_then(|v| v.to_str().ok())
197 .unwrap_or("")
198 .to_ascii_lowercase();
199 let req = if content_type.starts_with("application/json") {
200 match serde_json::from_slice::<CompleteRequest>(&body) {
201 Ok(r) => r,
202 Err(e) => {
203 return error_response(StatusCode::BAD_REQUEST, "OAUTH_INVALID_INPUT", format!("malformed JSON body: {e}"))
204 }
205 }
206 } else {
207 match parse_form_complete(&body) {
208 Some(r) => r,
209 None => {
210 return error_response(
211 StatusCode::BAD_REQUEST,
212 "OAUTH_INVALID_INPUT",
213 "body must be the callback form (code + state) or JSON with those fields",
214 )
215 }
216 }
217 };
218 match service.complete(principal.company_id, req).await {
219 Ok(out) => (StatusCode::OK, Json(serde_json::json!({ "success": true, "data": out }))).into_response(),
220 Err(e) => e.into_response(),
221 }
222}
223
224async fn disconnect(
226 State(service): State<Arc<IntegrationsOauthService>>,
227 axum::Extension(principal): axum::Extension<OAuthPrincipal>,
228 Path(account_id): Path<Uuid>,
229) -> Response {
230 if !principal.has("delete:integrations") {
231 return forbidden("delete:integrations");
232 }
233 match service.disconnect(principal.company_id, account_id).await {
234 Ok(()) => StatusCode::NO_CONTENT.into_response(),
235 Err(e) => e.into_response(),
236 }
237}
238
239async fn status(
243 State(service): State<Arc<IntegrationsOauthService>>,
244 axum::Extension(_principal): axum::Extension<OAuthPrincipal>,
245 Path(account_id): Path<Uuid>,
246) -> Response {
247 match service.status(account_id).await {
248 Ok(out) => (StatusCode::OK, Json(serde_json::json!({ "success": true, "data": out }))).into_response(),
249 Err(e) => e.into_response(),
250 }
251}
252
253pub fn create_oauth_routes(service: Arc<IntegrationsOauthService>) -> Router {
261 let public = Router::new().route("/oauth/callback", get(callback)).with_state(service.clone());
262 let authed = Router::new()
263 .route("/oauth/authorize", post(authorize))
264 .route("/oauth/complete", post(complete))
265 .route("/oauth/:id/disconnect", post(disconnect))
266 .route("/oauth/:id/status", get(status))
267 .layer(middleware::from_fn(require_principal))
268 .with_state(service);
269 public.merge(authed)
270}
271
272fn parse_form_complete(body: &[u8]) -> Option<CompleteRequest> {
281 let text = std::str::from_utf8(body).ok()?;
282 let mut code: Option<String> = None;
283 let mut state: Option<String> = None;
284 for pair in text.split('&') {
285 let (key, value) = pair.split_once('=')?;
286 let decoded = percent_decode(value);
287 match key {
288 "code" => code = Some(decoded),
289 "state" => state = Some(decoded),
290 _ => {}
291 }
292 }
293 let code = code.filter(|c| !c.is_empty())?;
294 let state = state.filter(|s| !s.is_empty())?;
295 Some(CompleteRequest { code, state })
296}
297
298fn percent_decode(input: &str) -> String {
299 let bytes = input.as_bytes();
300 let mut out = Vec::with_capacity(bytes.len());
301 let mut i = 0;
302 while i < bytes.len() {
303 match bytes[i] {
304 b'+' => {
305 out.push(b' ');
306 i += 1;
307 }
308 b'%' => {
309 if let Some(byte) = bytes.get(i + 1..i + 3).and_then(decode_hex_pair) {
310 out.push(byte);
311 i += 3;
312 } else {
313 out.push(b'%');
315 i += 1;
316 }
317 }
318 b => {
319 out.push(b);
320 i += 1;
321 }
322 }
323 }
324 String::from_utf8_lossy(&out).into_owned()
325}
326
327fn decode_hex_pair(pair: &[u8]) -> Option<u8> {
328 let hi = (pair.first()?).to_ascii_uppercase();
329 let lo = (pair.get(1)?).to_ascii_uppercase();
330 let digit = |b: u8| -> Option<u8> {
331 match b {
332 b'0'..=b'9' => Some(b - b'0'),
333 b'A'..=b'F' => Some(b - b'A' + 10),
334 _ => None,
335 }
336 };
337 Some(digit(hi)? * 16 + digit(lo)?)
338}
339
340fn error_page(status: StatusCode, title: &str, detail: &str) -> Response {
343 let escaped = detail
344 .replace('&', "&")
345 .replace('<', "<")
346 .replace('>', ">");
347 (
348 status,
349 Html(format!(
350 "<!doctype html>\n<html>\n<head><meta charset=\"utf-8\"><title>Connection not completed</title></head>\n<body>\n <h1>{title}</h1>\n <p>{escaped}</p>\n <p>Close this window and start the connection again.</p>\n</body>\n</html>\n"
351 )),
352 )
353 .into_response()
354}
355
356#[cfg(test)]
357mod tests {
358 use super::*;
359
360 #[test]
361 fn form_body_parses_code_and_state() {
362 let req = parse_form_complete(b"code=4%2F0Ax4PKb&state=abc.def").unwrap();
363 assert_eq!(req.code, "4/0Ax4PKb");
364 assert_eq!(req.state, "abc.def");
365 assert_eq!(percent_decode("a+b"), "a b");
367 assert!(parse_form_complete(b"code=only").is_none());
368 assert!(parse_form_complete(b"code=&state=x").is_none());
369 assert!(parse_form_complete(b"state=x").is_none());
370 }
371
372 #[test]
373 fn percent_decode_handles_odd_input_without_panicking() {
374 assert_eq!(percent_decode("100%"), "100%");
375 assert_eq!(percent_decode("%2z"), "%2z");
376 assert_eq!(percent_decode("%41%42"), "AB");
377 assert_eq!(percent_decode("caf%C3%A9"), "café");
378 }
379}