use std::collections::{HashMap, HashSet};
use std::sync::{Arc, Mutex};
use http::Uri;
use nio_client::session::token_hash;
use nio_client::wire;
use nio_client::{User, UserId};
use rand::RngCore;
use tokio::net::TcpListener;
use tokio_stream::wrappers::TcpListenerStream;
use tonic::{Request, Response, Status};
const SESSION_TTL_SECONDS: i64 = 3600;
struct Session {
principal: UserId,
expires_at_unix: i64,
}
#[derive(PartialEq, Eq, Hash)]
struct Grant {
ns: String,
obj: String,
rel: String,
subject: User,
}
#[derive(Default)]
struct State {
sessions: HashMap<String, Session>,
tuples: HashSet<Grant>,
names: HashMap<UserId, String>,
}
#[derive(Clone, Default)]
pub struct Backend {
state: Arc<Mutex<State>>,
}
impl Backend {
pub fn new() -> Self {
Backend::default()
}
fn lock(&self) -> std::sync::MutexGuard<'_, State> {
self.state.lock().expect("backend state poisoned")
}
pub fn register_principal(&self, principal: UserId, name: &str) {
self.lock().names.insert(principal, name.to_string());
}
pub fn grant(&self, ns: &str, obj: &str, rel: &str, subject: User) {
self.lock().tuples.insert(Grant {
ns: ns.to_string(),
obj: obj.to_string(),
rel: rel.to_string(),
subject,
});
}
pub fn create_session(&self, principal: UserId) -> String {
let mut bytes = [0u8; 32];
rand::thread_rng().fill_bytes(&mut bytes);
let raw = hex::encode(bytes);
let expires_at_unix = now_unix() + SESSION_TTL_SECONDS;
self.lock().sessions.insert(
token_hash(&raw),
Session {
principal,
expires_at_unix,
},
);
raw
}
pub fn remove_session(&self, hash: &str) {
self.lock().sessions.remove(hash);
}
pub async fn serve(self) -> Uri {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("bind nio backend");
let addr = listener.local_addr().expect("backend addr");
tokio::spawn(async move {
tonic::transport::Server::builder()
.add_service(wire::check_service_server::CheckServiceServer::new(
self.clone(),
))
.add_service(wire::session_service_server::SessionServiceServer::new(
self,
))
.serve_with_incoming(TcpListenerStream::new(listener))
.await
.expect("nio backend server");
});
format!("http://{addr}").parse().expect("backend uri")
}
}
fn now_unix() -> i64 {
chrono::Utc::now().timestamp()
}
fn short(hash: &str) -> &str {
&hash[..hash.len().min(8)]
}
#[tonic::async_trait]
impl wire::session_service_server::SessionService for Backend {
async fn resolve(
&self,
request: Request<wire::ResolveRequest>,
) -> Result<Response<wire::ResolveResponse>, Status> {
let hash = request.into_inner().token_hash;
let state = self.lock();
let outcome = match state.sessions.get(&hash) {
Some(s) if s.expires_at_unix > now_unix() => {
let name = state.names.get(&s.principal).cloned().unwrap_or_default();
println!(
"[nio session] resolve({}…) -> {} ({name})",
short(&hash),
s.principal
);
wire::resolve_response::Outcome::Session(wire::Session {
principal: s.principal.get(),
tenant_id: "demo".into(),
expires_at_unix_seconds: s.expires_at_unix,
})
}
_ => {
println!("[nio session] resolve({}…) -> NOT FOUND", short(&hash));
wire::resolve_response::Outcome::NotFound(wire::NotFound {})
}
};
Ok(Response::new(wire::ResolveResponse {
outcome: Some(outcome),
}))
}
}
#[tonic::async_trait]
impl wire::check_service_server::CheckService for Backend {
async fn check(
&self,
request: Request<wire::CheckRequest>,
) -> Result<Response<wire::CheckResponse>, Status> {
let req = request.into_inner();
let Some(wire::check_request::User::UserId(id)) = req.user else {
return Err(Status::invalid_argument(
"the demo only checks userId subjects",
));
};
let user_id = UserId::try_from(id).map_err(|e| Status::invalid_argument(e.to_string()))?;
let state = self.lock();
let key = |subject: User| Grant {
ns: req.ns.clone(),
obj: req.obj.clone(),
rel: req.rel.clone(),
subject,
};
let granted = state.tuples.contains(&key(user_id.into()));
let public = state.tuples.contains(&key(User::AllUsers));
let allowed = granted || public;
let verdict = if allowed { "ALLOW" } else { "DENY" };
let name = state
.names
.get(&user_id)
.cloned()
.unwrap_or_else(|| "?".into());
println!(
"[nio check] {}:{}#{} @ {} ({name}) -> {verdict}",
req.ns, req.obj, req.rel, user_id
);
Ok(Response::new(wire::CheckResponse {
principal: Some(wire::Principal { id: user_id.get() }),
ok: allowed,
}))
}
async fn content_change_check(
&self,
_request: Request<wire::ContentChangeCheckRequest>,
) -> Result<Response<wire::ContentChangeCheckResponse>, Status> {
Err(Status::unimplemented("not used in the webapp example"))
}
async fn list(
&self,
_request: Request<wire::ListRequest>,
) -> Result<Response<wire::ListResponse>, Status> {
Err(Status::unimplemented("not used in the webapp example"))
}
async fn expand(
&self,
_request: Request<wire::ExpandRequest>,
) -> Result<Response<wire::ExpandResponse>, Status> {
Err(Status::unimplemented("not used in the webapp example"))
}
async fn read(
&self,
_request: Request<wire::ReadRequest>,
) -> Result<Response<wire::ReadResponse>, Status> {
Err(Status::unimplemented("not used in the webapp example"))
}
async fn write(
&self,
_request: Request<wire::WriteRequest>,
) -> Result<Response<wire::WriteResponse>, Status> {
Err(Status::unimplemented("not used in the webapp example"))
}
type WatchStream = futures::stream::Empty<Result<wire::WatchResponse, Status>>;
async fn watch(
&self,
_request: Request<wire::WatchRequest>,
) -> Result<Response<Self::WatchStream>, Status> {
Err(Status::unimplemented("not used in the webapp example"))
}
}