use std::sync::Arc;
use systemprompt_identifiers::{SessionId, UserId};
use systemprompt_traits::AnalyticsProvider;
#[derive(Debug, thiserror::Error)]
pub enum SessionAttestationError {
#[error("Session missing or revoked")]
Missing,
#[error("Session user mismatch")]
UserMismatch,
#[error("Failed to check session: {0}")]
Lookup(String),
}
pub async fn attest_session(
analytics_provider: &Arc<dyn AnalyticsProvider>,
session_id: &SessionId,
user_id: &UserId,
route_context: &str,
) -> Result<(), SessionAttestationError> {
let session = analytics_provider
.find_active_session_by_id(session_id)
.await
.map_err(|e| SessionAttestationError::Lookup(e.to_string()))?;
let Some(session) = session else {
tracing::info!(
session_id = %session_id.as_str(),
user_id = %user_id.as_str(),
route = %route_context,
"session attestation failed: session missing or revoked"
);
return Err(SessionAttestationError::Missing);
};
if let Some(session_user_id) = session.user_id.as_ref()
&& session_user_id.as_str() != user_id.as_str()
{
tracing::warn!(
session_id = %session_id.as_str(),
claimed_user_id = %user_id.as_str(),
session_user_id = %session_user_id.as_str(),
route = %route_context,
"session attestation failed: session user mismatch"
);
return Err(SessionAttestationError::UserMismatch);
}
Ok(())
}