use std::convert::Infallible;
use std::sync::Arc;
use bytes::Bytes;
use http_body_util::{BodyExt, Full};
use hyper::body::Incoming;
use hyper::header::{AUTHORIZATION, CONTENT_TYPE, WWW_AUTHENTICATE};
use hyper::server::conn::http1;
use hyper::service::service_fn;
use hyper::{Method, Request, Response, StatusCode};
use hyper_util::rt::TokioIo;
use oauth_resource_server::testing;
use oauth_resource_server::{
Credential, KeyNaming, OAuthConfig, OAuthValidator, PROTECTED_RESOURCE_METADATA_PREFIX,
TokenRejection, authenticate, refusal,
};
use tokio::net::{TcpListener, TcpStream};
const ISSUER: &str = "https://auth.example.com/";
fn bearer(value: &str) -> Option<&str> {
let (scheme, token) = value.split_once(' ')?;
scheme.eq_ignore_ascii_case("bearer").then(|| token.trim())
}
async fn handle(
request: Request<Incoming>,
oauth: Arc<OAuthValidator>,
) -> Result<Response<Full<Bytes>>, Infallible> {
let path = request.uri().path();
if path == oauth.metadata_path() || path == PROTECTED_RESOURCE_METADATA_PREFIX {
if request.method() != Method::GET {
return Ok(status_only(StatusCode::METHOD_NOT_ALLOWED));
}
let mut response = Response::new(Full::from(oauth.metadata().to_string()));
response
.headers_mut()
.insert(CONTENT_TYPE, "application/json".parse().unwrap());
return Ok(response);
}
let candidate = request
.headers()
.get(AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(bearer);
match authenticate(candidate, None, Some(&oauth)).await {
Ok(Credential::OAuth(token)) => Ok(Response::new(Full::from(format!(
"hello, {}\n",
token.principal.as_deref().unwrap_or("caller")
)))),
Ok(_) => Ok(status_only(StatusCode::INTERNAL_SERVER_ERROR)),
Err(rejection) => {
if rejection == TokenRejection::Missing {
tracing::debug!(path, "no credential presented");
} else {
tracing::warn!(path, reason = ?rejection, "refused");
}
let r = refusal(&rejection, Some(&oauth));
let mut response = status_only(StatusCode::from_u16(r.status).unwrap());
if let Some(challenge) = r.www_authenticate {
response
.headers_mut()
.insert(WWW_AUTHENTICATE, challenge.parse().unwrap());
}
Ok(response)
}
}
}
fn status_only(status: StatusCode) -> Response<Full<Bytes>> {
let mut response = Response::new(Full::default());
*response.status_mut() = status;
response
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
tracing_subscriber::fmt().init();
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let resolved = OAuthConfig {
enabled: true,
issuer: ISSUER.into(),
jwks_uri: Some(jwks.url.clone()),
audience: "example-api".into(),
resource: "https://api.example.com".into(),
required_scopes: vec!["api:read".into()],
..OAuthConfig::default()
}
.resolve(KeyNaming::Dotted("oauth"))?
.expect("enabled: true");
let oauth = Arc::new(OAuthValidator::new(&resolved)?);
oauth.spawn_background_refresh();
let listener = TcpListener::bind("127.0.0.1:0").await?;
let addr = listener.local_addr()?;
let server_oauth = Arc::clone(&oauth);
tokio::spawn(async move {
loop {
let Ok((stream, _)) = listener.accept().await else {
continue;
};
let oauth = Arc::clone(&server_oauth);
tokio::spawn(async move {
let service = service_fn(move |request| handle(request, Arc::clone(&oauth)));
if let Err(e) = http1::Builder::new()
.serve_connection(TokioIo::new(stream), service)
.await
{
tracing::debug!("connection error: {e}");
}
});
}
});
let token = |scope: &str| {
testing::mint(
testing::KEY_A_PEM,
testing::KID_A,
&serde_json::json!({
"iss": ISSUER, "aud": "example-api", "sub": "example-user",
"exp": testing::now() + 300, "scope": scope,
}),
)
};
let requests = [
("no credential", "/api", None),
(
"valid token",
"/api",
Some(format!("Bearer {}", token("api:read"))),
),
(
"missing scope",
"/api",
Some(format!("Bearer {}", token("other"))),
),
("not a JWT", "/api", Some("Bearer not-a-jwt".to_string())),
("metadata", "/.well-known/oauth-protected-resource", None),
];
for (label, path, authorization) in requests {
let (status, challenge, body) = send(addr, path, authorization.as_deref()).await?;
println!("{label}: {status}");
if let Some(challenge) = challenge {
println!(" WWW-Authenticate: {challenge}");
}
if !body.is_empty() {
println!(" body: {}", body.trim_end());
}
}
Ok(())
}
async fn send(
addr: std::net::SocketAddr,
path: &str,
authorization: Option<&str>,
) -> Result<(StatusCode, Option<String>, String), Box<dyn std::error::Error>> {
let stream = TcpStream::connect(addr).await?;
let (mut sender, connection) =
hyper::client::conn::http1::handshake(TokioIo::new(stream)).await?;
tokio::spawn(connection);
let mut request = Request::builder()
.uri(path)
.header("host", addr.to_string());
if let Some(value) = authorization {
request = request.header(AUTHORIZATION, value);
}
let response = sender
.send_request(request.body(Full::<Bytes>::default())?)
.await?;
let status = response.status();
let challenge = response
.headers()
.get(WWW_AUTHENTICATE)
.map(|v| v.to_str().unwrap_or_default().to_string());
let body = response.into_body().collect().await?.to_bytes();
Ok((
status,
challenge,
String::from_utf8_lossy(&body).into_owned(),
))
}