use std::sync::Arc;
use axum::{
Extension, Json, Router,
extract::{Path, Query, State},
http::StatusCode,
response::{IntoResponse, Response},
routing::get,
};
use fraiseql_auth::saml::{SamlError, SamlIdpRecord, SamlIdpRegistry, SamlIdpSpec};
use serde::{Deserialize, Serialize};
use uuid::Uuid;
use super::admin_principal::{AdminPrincipal, foreign_tenant_response};
#[derive(Clone)]
pub struct SamlIdpManagementState {
pub registry: SamlIdpRegistry,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SamlIdpDto {
pub id: Uuid,
pub idp_name: String,
pub tenant_id: Option<Uuid>,
pub sp_entity_id: String,
pub acs_url: String,
pub metadata_xml: String,
pub idp_entity_id: String,
pub trust_asserted_email: bool,
pub certificate_expires_at: Option<chrono::DateTime<chrono::Utc>>,
pub created_at: chrono::DateTime<chrono::Utc>,
pub updated_at: chrono::DateTime<chrono::Utc>,
}
impl From<SamlIdpRecord> for SamlIdpDto {
fn from(r: SamlIdpRecord) -> Self {
Self {
id: r.id,
idp_name: r.idp_name,
tenant_id: r.tenant_id,
sp_entity_id: r.sp_entity_id,
acs_url: r.acs_url,
metadata_xml: r.metadata_xml,
idp_entity_id: r.idp_entity_id,
trust_asserted_email: r.trust_asserted_email,
certificate_expires_at: r.certificate_expires_at,
created_at: r.created_at,
updated_at: r.updated_at,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct CreateSamlIdpRequest {
pub idp_name: String,
#[serde(default)]
pub tenant_id: Option<Uuid>,
pub sp_entity_id: String,
pub acs_url: String,
pub metadata_xml: String,
#[serde(default)]
pub trust_asserted_email: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct UpdateSamlIdpRequest {
pub sp_entity_id: String,
pub acs_url: String,
pub metadata_xml: String,
#[serde(default)]
pub trust_asserted_email: bool,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ListQuery {
#[serde(default)]
pub tenant_id: Option<Uuid>,
}
fn json_error(status: StatusCode, message: &str) -> Response {
(status, Json(serde_json::json!({ "error": message }))).into_response()
}
fn store_error(e: &SamlError) -> Response {
match e {
SamlError::NotFound(_) => json_error(StatusCode::NOT_FOUND, &e.to_string()),
SamlError::NameTaken(_) => json_error(StatusCode::CONFLICT, &e.to_string()),
SamlError::Config(_) => json_error(StatusCode::BAD_REQUEST, &e.to_string()),
_ => {
tracing::error!(error = %e, "SAML IdP management operation failed");
json_error(StatusCode::INTERNAL_SERVER_ERROR, "SAML IdP store error")
},
}
}
pub fn saml_idp_management_router(state: SamlIdpManagementState) -> Router {
Router::new()
.route("/api/saml/idps", axum::routing::post(create_idp).get(list_idps))
.route("/api/saml/idps/{idp_name}", get(get_idp).put(update_idp).delete(delete_idp))
.with_state(Arc::new(state))
}
async fn create_idp(
State(state): State<Arc<SamlIdpManagementState>>,
Extension(principal): Extension<AdminPrincipal>,
Json(payload): Json<CreateSamlIdpRequest>,
) -> Response {
let Ok(tenant_id) = principal.scope(payload.tenant_id) else {
return foreign_tenant_response();
};
let spec = SamlIdpSpec {
idp_name: payload.idp_name,
tenant_id,
sp_entity_id: payload.sp_entity_id,
acs_url: payload.acs_url,
metadata_xml: payload.metadata_xml,
trust_asserted_email: payload.trust_asserted_email,
};
match state.registry.create(&spec).await {
Ok(record) => (StatusCode::CREATED, Json(SamlIdpDto::from(record))).into_response(),
Err(e) => store_error(&e),
}
}
async fn list_idps(
State(state): State<Arc<SamlIdpManagementState>>,
Extension(principal): Extension<AdminPrincipal>,
Query(q): Query<ListQuery>,
) -> Response {
let Ok(tenant) = principal.scope(q.tenant_id) else {
return foreign_tenant_response();
};
match state.registry.list_stored().await {
Ok(records) => {
let idps: Vec<SamlIdpDto> = records
.into_iter()
.filter(|r| tenant.is_none_or(|t| r.tenant_id == Some(t)))
.map(SamlIdpDto::from)
.collect();
Json(serde_json::json!({ "total": idps.len(), "idps": idps })).into_response()
},
Err(e) => store_error(&e),
}
}
async fn get_idp(
State(state): State<Arc<SamlIdpManagementState>>,
Extension(principal): Extension<AdminPrincipal>,
Path(idp_name): Path<String>,
) -> Response {
match state.registry.get_stored(&idp_name).await {
Ok(Some(record)) if principal.may_see(record.tenant_id) => {
Json(SamlIdpDto::from(record)).into_response()
},
Ok(_) => json_error(StatusCode::NOT_FOUND, "no such SAML IdP"),
Err(e) => store_error(&e),
}
}
async fn update_idp(
State(state): State<Arc<SamlIdpManagementState>>,
Extension(principal): Extension<AdminPrincipal>,
Path(idp_name): Path<String>,
Json(payload): Json<UpdateSamlIdpRequest>,
) -> Response {
let existing = match state.registry.get_stored(&idp_name).await {
Ok(Some(record)) if principal.may_see(record.tenant_id) => record,
Ok(_) => return json_error(StatusCode::NOT_FOUND, "no such SAML IdP"),
Err(e) => return store_error(&e),
};
let spec = SamlIdpSpec {
idp_name,
tenant_id: existing.tenant_id,
sp_entity_id: payload.sp_entity_id,
acs_url: payload.acs_url,
metadata_xml: payload.metadata_xml,
trust_asserted_email: payload.trust_asserted_email,
};
match state.registry.update(&spec).await {
Ok(record) => Json(SamlIdpDto::from(record)).into_response(),
Err(e) => store_error(&e),
}
}
async fn delete_idp(
State(state): State<Arc<SamlIdpManagementState>>,
Extension(principal): Extension<AdminPrincipal>,
Path(idp_name): Path<String>,
) -> Response {
match state.registry.get_stored(&idp_name).await {
Ok(Some(record)) if principal.may_see(record.tenant_id) => {},
Ok(_) => return json_error(StatusCode::NOT_FOUND, "no such SAML IdP"),
Err(e) => return store_error(&e),
}
match state.registry.delete(&idp_name).await {
Ok(()) => StatusCode::NO_CONTENT.into_response(),
Err(e) => store_error(&e),
}
}
#[cfg(test)]
mod tests;