use akari::Value;
use hotaru_core::app::common::{RunMode, RuntimeConfig};
use hotaru_core::connection::TransportSpec;
use hotaru_core::connection::error::ConnectionError;
use hotaru_core::debug_log;
use hotaru_core::extensions::{Locals, Params};
use hotaru_core::protocol::{
BoxProtocolError, EndpointOutcome, ProtocolError, ProtocolRole, RequestContext,
};
use hotaru_core::url::UrlNode;
use hotaru_core::connection::{HotaruBufRead, HotaruWrite};
use once_cell::sync::Lazy;
use std::collections::HashMap;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use crate::channel::Http1Channel;
use crate::message::body::HttpBody;
use crate::message::http_value::{HttpMethod, StatusCode};
use crate::message::meta::HttpMeta;
use crate::message::request::HttpRequest;
use crate::message::response::{HttpResponse, response_templates};
use crate::protocol::HttpError;
use crate::security::safety::HttpSafety;
use crate::util::cookie::{Cookie, CookieMap};
use crate::util::form::{MultiForm, UrlEncodedForm};
pub enum Executable<TS: TransportSpec = hotaru_io_tokio::TcpTransport> {
Request {
runtime: Arc<RuntimeConfig>,
endpoint: Arc<UrlNode<HttpContext<TS>, TS>>,
},
Response,
}
pub struct HttpContext<TS: TransportSpec = hotaru_io_tokio::TcpTransport> {
pub request: HttpRequest,
pub response: HttpResponse,
pub executable: Executable<TS>,
pub host: Option<String>, pub safety: HttpSafety,
remote_addr: Option<SocketAddr>,
local_addr: Option<SocketAddr>,
pub params: Params,
pub locals: Locals,
channel: Option<Http1Channel<TS::Wire>>,
}
pub type HttpReqCtx<TS = hotaru_io_tokio::TcpTransport> = HttpContext<TS>;
const UNSET_ADDR: SocketAddr =
SocketAddr::new(std::net::IpAddr::V4(std::net::Ipv4Addr::new(0, 0, 0, 0)), 0);
impl<TS: TransportSpec> HttpContext<TS> {
pub fn new_server(
runtime: Arc<RuntimeConfig>,
endpoint: Arc<UrlNode<HttpContext<TS>, TS>>,
request: HttpRequest,
remote_addr: Option<SocketAddr>,
local_addr: Option<SocketAddr>,
safety: HttpSafety,
) -> Self {
Self {
request,
response: HttpResponse::default(),
executable: Executable::Request { runtime, endpoint },
host: None,
safety,
remote_addr,
local_addr,
params: Default::default(),
locals: Default::default(),
channel: None,
}
}
pub fn new_client(host: String, safety: HttpSafety) -> Self {
Self {
request: HttpRequest::default(),
response: HttpResponse::default(),
executable: Executable::<TS>::Response,
host: if host.is_empty() { None } else { Some(host) },
safety,
remote_addr: None,
local_addr: None,
params: Default::default(),
locals: Default::default(),
channel: None,
}
}
pub(crate) fn install_channel(&mut self, channel: Http1Channel<TS::Wire>) {
self.channel = Some(channel);
}
pub(crate) fn channel(&self) -> Option<&Http1Channel<TS::Wire>> {
self.channel.as_ref()
}
#[inline]
pub fn client_ip(&self) -> Option<SocketAddr> {
self.remote_addr
}
#[inline]
pub fn client_ip_or_default(&self) -> SocketAddr {
self.remote_addr.unwrap_or(UNSET_ADDR)
}
#[inline]
pub fn client_ip_only(&self) -> Option<IpAddr> {
self.remote_addr.map(|addr| addr.ip())
}
#[inline]
pub fn client_ip_only_or_default(&self) -> IpAddr {
self.remote_addr
.map(|addr| addr.ip())
.unwrap_or(UNSET_ADDR.ip())
}
#[inline]
pub fn server_addr(&self) -> Option<SocketAddr> {
self.local_addr
}
#[inline]
pub fn server_addr_or_default(&self) -> SocketAddr {
self.local_addr.unwrap_or(UNSET_ADDR)
}
#[inline]
pub fn remote_addr(&self) -> Option<SocketAddr> {
self.remote_addr
}
#[inline]
pub fn remote_addr_or_default(&self) -> SocketAddr {
self.remote_addr.unwrap_or(UNSET_ADDR)
}
#[inline]
pub fn local_addr(&self) -> Option<SocketAddr> {
self.local_addr
}
#[inline]
pub fn local_addr_or_default(&self) -> SocketAddr {
self.local_addr.unwrap_or(UNSET_ADDR)
}
pub async fn read_request<R>(
runtime: Arc<RuntimeConfig>,
reader: &mut R,
) -> Result<HttpRequest, ConnectionError>
where
R: HotaruBufRead<Error = std::io::Error> + Unpin + Send,
{
Ok(HttpRequest::parse_lazy(
reader,
&runtime.get_config::<HttpSafety>().unwrap_or_default(),
runtime.mode() == RunMode::Build,
)
.await)
}
pub async fn send_response<W>(response: HttpResponse, writer: &mut W)
where
W: HotaruWrite<Error = std::io::Error> + Unpin + Send,
{
let _ = response.send(writer).await;
}
pub async fn run(mut self) -> Result<HttpContext<TS>, BoxProtocolError> {
if let Some(endpoint) = self.endpoint() {
debug_log!("HTTP Context: Found endpoint, checking request");
if let Err(err) = self.request_check(&endpoint) {
debug_log!("HTTP Context: Request check failed: {:?}", err);
let status: StatusCode = (&err).into();
self.response = response_templates::return_status(status);
return Ok(self);
};
debug_log!("HTTP Context: Running endpoint handler");
let result = endpoint.run(self).await.map_err(|e| e.boxed());
debug_log!("HTTP Context: Handler completed");
result
} else {
debug_log!("HTTP Context: No endpoint available (client context)");
Ok(self)
}
}
pub fn request_check(
&mut self,
endpoint: &Arc<UrlNode<HttpContext<TS>, TS>>,
) -> Result<(), HttpError> {
let mut config = self.safety.clone();
if let Some(ep) = endpoint.get_params::<HttpSafety>() {
config.update(&ep);
}
if !config.check_body_size(self.request.meta.get_content_length().unwrap_or(0)) {
return Err(HttpError::PayloadTooLarge);
}
if !config.check_method(&self.request.meta.method()) {
return Err(HttpError::MethodNotAllowed);
}
if !config.check_content_type(&self.request.meta.get_content_type().unwrap_or_default()) {
return Err(HttpError::UnsupportedMediaType);
}
return Ok(());
}
pub fn meta(&mut self) -> &mut HttpMeta {
&mut self.request.meta
}
pub fn runtime(&self) -> Option<Arc<RuntimeConfig>> {
match &self.executable {
Executable::Request { runtime, .. } => Some(runtime.clone()),
Executable::<TS>::Response => None,
}
}
pub fn endpoint(&self) -> Option<Arc<UrlNode<HttpContext<TS>, TS>>> {
match &self.executable {
Executable::Request { endpoint, .. } => Some(endpoint.clone()),
Executable::<TS>::Response => None,
}
}
pub async fn parse_body(&mut self) {
let mut settings = self.safety.clone();
if let Some(endpoint) = self.endpoint() {
if let Some(ep) = endpoint.get_params::<HttpSafety>() {
settings.update(&ep);
}
}
let body = std::mem::take(&mut self.request.body);
self.request.body = body.parse_buffer(&settings);
}
pub async fn form(&mut self) -> Option<&UrlEncodedForm> {
self.parse_body().await; if let HttpBody::Form(ref data) = self.request.body {
Some(data)
} else {
None
}
}
pub async fn form_or_default(&mut self) -> &UrlEncodedForm {
match self.form().await {
Some(form) => form,
None => {
static EMPTY: Lazy<UrlEncodedForm> = Lazy::new(|| HashMap::new().into());
&EMPTY
}
}
}
pub async fn files(&mut self) -> Option<&MultiForm> {
self.parse_body().await; if let HttpBody::Files(ref data) = self.request.body {
Some(data)
} else {
None
}
}
pub async fn files_or_default(&mut self) -> &MultiForm {
match self.files().await {
Some(files) => files,
None => {
static EMPTY: Lazy<MultiForm> = Lazy::new(|| HashMap::new().into());
&EMPTY
}
}
}
pub async fn json(&mut self) -> Option<&Value> {
self.parse_body().await; if let HttpBody::Json(ref data) = self.request.body {
Some(data)
} else {
None
}
}
pub async fn json_or_default(&mut self) -> &Value {
match self.json().await {
Some(json) => json,
None => {
static EMPTY: Lazy<Value> = Lazy::new(|| Value::new(""));
&EMPTY
}
}
}
pub fn segment(&mut self, index: usize) -> String {
self.request.meta.get_path(index + 1)
}
pub fn path(&self) -> String {
self.request.meta.path()
}
pub fn param<A: AsRef<str>>(&mut self, name: A) -> Option<String> {
self.endpoint().and_then(|endpoint| {
endpoint
.match_seg_name_with_index(name)
.map(|index| self.request.meta.get_path(index))
})
}
pub fn pattern<A: AsRef<str>>(&mut self, name: A) -> Option<String> {
self.param(name)
}
pub fn query<T: Into<String>>(&mut self, key: T) -> Option<String> {
self.request.meta.get_url_args(key)
}
pub fn get_preferred_language(&mut self) -> Option<String> {
self.request
.meta
.get_lang()
.map(|lang_dict| lang_dict.most_preferred())
}
pub fn get_preferred_language_or_default<T: AsRef<str>>(&mut self, default: T) -> String {
self.get_preferred_language()
.unwrap_or_else(|| default.as_ref().to_string())
}
pub fn method(&mut self) -> HttpMethod {
self.request.meta.method()
}
pub fn headers(&self) -> &HashMap<String, crate::message::meta::HeaderValue> {
&self.request.meta.header
}
pub fn header(&self, key: &str) -> Option<&crate::message::meta::HeaderValue> {
self.request.meta.header.get(key)
}
pub fn header_str(&self, key: &str) -> Option<&str> {
self.request.meta.header.get(key).and_then(|hv| match hv {
crate::message::meta::HeaderValue::Single(s) => Some(s.as_str()),
crate::message::meta::HeaderValue::Multiple(v) => v.first().map(|s| s.as_str()),
})
}
pub fn has_header(&self, key: &str) -> bool {
self.request.meta.header.contains_key(key)
}
pub fn get_cookies(&mut self) -> &CookieMap {
self.request.meta.get_cookies()
}
pub fn get_cookie(&mut self, key: &str) -> Option<Cookie> {
self.request.meta.get_cookie(key)
}
pub fn get_cookie_or_default<T: AsRef<str>>(&mut self, key: T) -> Cookie {
self.request.meta.get_cookie_or_default(key)
}
pub fn response_mut(&mut self) -> &mut HttpResponse {
&mut self.response
}
pub fn set_status(&mut self, code: u16) -> &mut Self {
self.response.meta.start_line.set_status_code(code);
self
}
pub fn add_response_header(&mut self, key: String, value: String) -> &mut Self {
self.response.meta.set_attribute(key, value);
self
}
pub fn set_body(&mut self, body: HttpBody) -> &mut Self {
self.response.body = body;
self
}
pub(crate) fn take_request(&mut self) -> HttpRequest {
if self.request.meta.get_host().is_none() {
if let Some(host) = self.host.as_deref().filter(|h| !h.is_empty()) {
self.request.meta.set_host(Some(host.to_string()));
}
}
std::mem::take(&mut self.request)
}
pub(crate) fn set_response(&mut self, response: HttpResponse) {
self.response = response;
}
}
impl<TS: TransportSpec> RequestContext for HttpContext<TS> {
type Request = HttpRequest;
type Response = HttpResponse;
type Error = crate::protocol::HttpError;
type Channel = Http1Channel<TS::Wire>;
fn handle_error(&mut self) {
match &self.executable {
Executable::Request { .. } => {
self.response = response_templates::html_response(
"<h1>500 Internal Server Error</h1><br><p>An unexpected error occurred</p>",
)
.status(500);
}
Executable::<TS>::Response => {
self.response = HttpResponse::default();
}
}
}
fn role(&self) -> ProtocolRole {
match &self.executable {
Executable::Request { .. } => ProtocolRole::Server,
Executable::<TS>::Response => ProtocolRole::Client,
}
}
fn inject_request(&mut self, request: Self::Request) {
self.request(request);
}
fn into_response(self) -> Self::Response {
self.response
}
}
impl<TS: TransportSpec> EndpointOutcome<HttpContext<TS>> for HttpResponse {
fn apply_to(self, ctx: &mut HttpContext<TS>) -> Result<(), HttpError> {
ctx.response = self;
Ok(())
}
}
impl<TS: TransportSpec> Default for HttpContext<TS> {
fn default() -> Self {
Self::new_client(String::new(), HttpSafety::default())
}
}
impl<TS: TransportSpec> HttpContext<TS> {
pub fn bad_request(&mut self) {
self.handle_error();
}
}
pub type HttpResCtx<TS = hotaru_io_tokio::TcpTransport> = HttpContext<TS>;
impl<TS: TransportSpec> HttpContext<TS> {
pub fn new_res(config: HttpSafety, host: impl Into<String>) -> Self {
Self::new_client(host.into(), config)
}
pub fn request(&mut self, mut request: HttpRequest) {
if request.meta.get_host().is_none() {
if let Some(host) = self.host.as_deref().filter(|h| !h.is_empty()) {
request.meta.set_host(Some(host.to_string()));
}
};
self.request = request;
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::message::http_value::StatusCode;
use crate::message::response::response_templates;
type TestHttpContext = HttpContext<hotaru_io_tokio::TcpTransport>;
fn client_context(host: &str) -> TestHttpContext {
TestHttpContext::new_client(host.to_string(), HttpSafety::default())
}
#[test]
fn take_request_sets_missing_host_from_context() {
let mut ctx = client_context("example.com");
let mut request = ctx.take_request();
assert_eq!(request.meta.get_host(), Some("example.com".to_string()));
}
#[test]
fn take_request_preserves_existing_request_host() {
let mut ctx = client_context("context.example");
let mut request = HttpRequest::default();
request.meta.set_host(Some("request.example".to_string()));
ctx.request = request;
let mut request = ctx.take_request();
assert_eq!(request.meta.get_host(), Some("request.example".to_string()));
}
#[test]
fn take_request_ignores_empty_context_host() {
let mut ctx = client_context("");
let mut request = ctx.take_request();
assert_eq!(request.meta.get_host(), None);
}
#[test]
fn set_response_stores_response() {
let mut ctx = client_context("");
let response = response_templates::normal_response(StatusCode::CREATED, "created");
ctx.set_response(response);
assert_eq!(
ctx.response.meta.start_line.status_code(),
StatusCode::CREATED
);
}
}