use std::collections::HashMap;
use std::fmt::Debug;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::time::Duration;
use thiserror::Error;
use tokio::sync::{broadcast, RwLock};
use tokio::task::JoinHandle;
use tracing::{debug, error, info};
use crate::http_tcp::{HttpTcpRequest, HttpTcpResponse, not_found};
use crate::message::Message;
use crate::registration::{
Capabilities, Endpoint, HealthStatus, RegisteredModule, advertise_capabilities,
start_heartbeat_loop,
};
use crate::tcp_channel::TcpChannel;
use crate::communication::MessageChannel;
use crate::tcp_types::ConnectionConfig;
#[derive(Debug, Error)]
pub enum HttpTcpError {
#[error("Invalid request: {0}")]
InvalidRequest(String),
#[error("Not found: {0}")]
NotFound(String),
#[error("Unauthorized: {0}")]
Unauthorized(String),
#[error("Forbidden: {0}")]
Forbidden(String),
#[error("Internal error: {0}")]
Internal(String),
#[error("SDK error: {0}")]
Sdk(#[from] Box<crate::error::Error>),
#[error("JSON error: {0}")]
Json(#[from] serde_json::Error),
#[error("I/O error: {0}")]
Io(#[from] std::io::Error),
#[error("Timeout after {0:?}")]
Timeout(Duration),
#[error("Other error: {0}")]
Other(String),
}
pub type HttpTcpResult<T> = std::result::Result<T, HttpTcpError>;
pub type HandlerFn<S> = Arc<
dyn Fn(
HttpTcpRequest,
Arc<S>,
) -> Pin<Box<dyn Future<Output = HttpTcpResult<HttpTcpResponse>> + Send>>
+ Send
+ Sync,
>;
#[derive(Clone)]
pub struct Route<S> {
pub method: String,
pub path: String,
pub handler: HandlerFn<S>,
}
#[derive(Clone)]
pub struct HttpTcpRouter<S: Send + Sync + Clone + 'static> {
routes: Arc<RwLock<Vec<Route<S>>>>,
not_found_handler: Arc<RwLock<Option<HandlerFn<S>>>>,
}
impl<S: Send + Sync + Clone + 'static> Default for HttpTcpRouter<S> {
fn default() -> Self {
Self::new()
}
}
impl<S: Send + Sync + Clone + 'static> HttpTcpRouter<S> {
pub fn new() -> Self {
Self {
routes: Arc::new(RwLock::new(Vec::new())),
not_found_handler: Arc::new(RwLock::new(None)),
}
}
pub fn route<F, Fut>(self, method: &str, path: &str, handler: F) -> Self
where
F: Fn(HttpTcpRequest, Arc<S>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = HttpTcpResult<HttpTcpResponse>> + Send + 'static,
{
let handler = Arc::new(move |req, state| {
let fut = handler(req, state);
Box::pin(fut) as Pin<Box<dyn Future<Output = HttpTcpResult<HttpTcpResponse>> + Send>>
});
let route = Route {
method: method.to_uppercase(),
path: path.to_string(),
handler,
};
let routes_clone = self.routes.clone();
tokio::spawn(async move {
let mut routes = routes_clone.write().await;
routes.push(route);
});
self
}
pub fn not_found_handler<F, Fut>(self, handler: F) -> Self
where
F: Fn(HttpTcpRequest, Arc<S>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = HttpTcpResult<HttpTcpResponse>> + Send + 'static,
{
let handler = Arc::new(move |req, state| {
let fut = handler(req, state);
Box::pin(fut) as Pin<Box<dyn Future<Output = HttpTcpResult<HttpTcpResponse>> + Send>>
});
let not_found_handler_clone = self.not_found_handler.clone();
tokio::spawn(async move {
let mut not_found = not_found_handler_clone.write().await;
*not_found = Some(handler);
});
self
}
pub fn methods<F, Fut>(self, methods: Vec<&str>, path: &str, handler: F) -> Self
where
F: Fn(HttpTcpRequest, Arc<S>) -> Fut + Send + Sync + Clone + 'static,
Fut: Future<Output = HttpTcpResult<HttpTcpResponse>> + Send + 'static,
{
let mut result = self;
for method in methods {
let handler_clone = handler.clone();
result = result.route(method, path, handler_clone);
}
result
}
pub async fn handle_request(&self, request: HttpTcpRequest, state: Arc<S>) -> HttpTcpResponse {
let method = request.method.to_uppercase();
let path = request.path();
debug!("Routing request: {} {}", method, path);
let routes = self.routes.read().await;
for route in routes.iter() {
if route.method == method && route.path == path {
debug!("Found matching route: {} {}", route.method, route.path);
match (route.handler)(request.clone(), state.clone()).await {
Ok(response) => {
return response;
}
Err(e) => {
error!("Handler error: {}", e);
return error_to_response(e, &request.request_id);
}
}
}
}
let not_found_handler = self.not_found_handler.read().await;
if let Some(handler) = not_found_handler.as_ref() {
match handler(request.clone(), state).await {
Ok(response) => {
response
}
Err(e) => {
error!("Not found handler error: {}", e);
error_to_response(e, &request.request_id)
}
}
} else {
not_found(&request.request_id, path)
}
}
pub async fn get_routes(&self) -> Vec<Endpoint> {
let routes = self.routes.read().await;
let mut path_methods: HashMap<String, Vec<String>> = HashMap::new();
for route in routes.iter() {
let entry = path_methods.entry(route.path.clone()).or_default();
entry.push(route.method.clone());
}
path_methods
.into_iter()
.map(|(path, methods)| Endpoint::new(path, methods.iter().map(|s| s.as_str()).collect()))
.collect()
}
}
fn error_to_response(error: HttpTcpError, request_id: &str) -> HttpTcpResponse {
let status_code = match &error {
HttpTcpError::InvalidRequest(_) => 400,
HttpTcpError::NotFound(_) => 404,
HttpTcpError::Unauthorized(_) => 401,
HttpTcpError::Forbidden(_) => 403,
HttpTcpError::Timeout(_) => 408,
_ => 500,
};
let mut headers = HashMap::new();
headers.insert("Content-Type".to_string(), "application/json".to_string());
match serde_json::to_vec(&crate::http_tcp::ApiResponse::<()> {
status: "error".to_string(),
data: None,
message: Some(error.to_string()),
}) {
Ok(body) => HttpTcpResponse {
request_id: request_id.to_string(),
status_code,
headers,
body: Some(body),
},
Err(_) => {
headers.insert("Content-Type".to_string(), "text/plain".to_string());
HttpTcpResponse {
request_id: request_id.to_string(),
status_code,
headers,
body: Some(error.to_string().into_bytes()),
}
}
}
}
#[derive(Debug, Clone)]
pub struct HttpReceiver {
pub channel: Arc<TcpChannel>,
pub module: RegisteredModule,
}
pub async fn start_http_server<S: Send + Sync + Clone + 'static>(
router: HttpTcpRouter<S>,
module: RegisteredModule,
state: S,
shutdown_signal: Option<broadcast::Receiver<()>>,
) -> crate::error::Result<JoinHandle<()>> {
let config = ConnectionConfig::new(
module.orchestrator_host.clone(),
module.orchestrator_port,
);
let channel = Arc::new(TcpChannel::connect(config).await?);
let endpoints = router.get_routes().await;
let capabilities = Capabilities {
endpoints,
message_types: None,
additional_capabilities: None,
};
advertise_capabilities(&module, capabilities).await?;
let _heartbeat_task = start_heartbeat_loop(
module.clone(),
Duration::from_secs(30),
|| (HealthStatus::Healthy, None),
);
let state = Arc::new(state);
let _receiver = HttpReceiver {
channel: channel.clone(),
module: module.clone(),
};
info!("Starting HTTP-over-TCP server for module {}", module.info.name);
let server_handle = if let Some(mut shutdown_signal) = shutdown_signal {
tokio::spawn(async move {
let router = router.clone();
let state = state.clone();
loop {
tokio::select! {
_ = shutdown_signal.recv() => {
info!("Shutdown signal received, stopping HTTP-over-TCP server");
break;
}
message_result = channel.receive() => {
match message_result {
Ok(encoded_message) => {
let request_json = match encoded_message.format() {
crate::message::EncodingFormat::Json => {
match std::str::from_utf8(encoded_message.data()) {
Ok(s) => s.to_string(),
Err(e) => {
error!("Invalid UTF-8: {}", e);
return;
}
}
},
_ => {
let json_encoded = match encoded_message.to_format(crate::message::EncodingFormat::Json) {
Ok(j) => j,
Err(e) => {
error!("Failed to convert message to JSON: {}", e);
return;
}
};
match std::str::from_utf8(json_encoded.data()) {
Ok(s) => s.to_string(),
Err(e) => {
error!("Invalid UTF-8: {}", e);
return;
}
}
}
};
let message: Message<HttpTcpRequest> = match serde_json::from_str(&request_json) {
Ok(msg) => msg,
Err(e) => {
error!("Failed to deserialize request: {}", e);
return;
}
};
let request = message.content();
info!("Received HTTP-over-TCP request: {} {}", request.method, request.uri);
let router_clone = router.clone();
let state_clone = state.clone();
let channel_clone = channel.clone();
let request_clone = request.clone();
tokio::spawn(async move {
let response = router_clone.handle_request(request_clone, state_clone).await;
let response_message = Message::new(response);
match response_message.encode() {
Ok(encoded) => {
if let Err(e) = MessageChannel::send(&*channel_clone, encoded).await {
error!("Failed to send HTTP-over-TCP response: {}", e);
}
}
Err(e) => {
error!("Failed to encode HTTP-over-TCP response: {}", e);
}
}
});
}
Err(e) => {
error!("Failed to receive HTTP-over-TCP request: {}", e);
}
}
}
}
}
info!("HTTP-over-TCP server stopped");
})
} else {
tokio::spawn(async move {
let router = router.clone();
let state = state.clone();
loop {
match MessageChannel::receive(&*channel).await {
Ok(encoded_message) => {
let request_json = match encoded_message.format() {
crate::message::EncodingFormat::Json => {
match std::str::from_utf8(encoded_message.data()) {
Ok(s) => s.to_string(),
Err(e) => {
error!("Invalid UTF-8: {}", e);
continue;
}
}
},
_ => {
let json_encoded = match encoded_message.to_format(crate::message::EncodingFormat::Json) {
Ok(j) => j,
Err(e) => {
error!("Failed to convert message to JSON: {}", e);
continue;
}
};
match std::str::from_utf8(json_encoded.data()) {
Ok(s) => s.to_string(),
Err(e) => {
error!("Invalid UTF-8: {}", e);
continue;
}
}
}
};
let message: Message<HttpTcpRequest> = match serde_json::from_str(&request_json) {
Ok(msg) => msg,
Err(e) => {
error!("Failed to deserialize request: {}", e);
continue;
}
};
let request = message.content();
info!("Received HTTP-over-TCP request: {} {}", request.method, request.uri);
let router_clone = router.clone();
let state_clone = state.clone();
let channel_clone = channel.clone();
let request_clone = request.clone();
tokio::spawn(async move {
let response = router_clone.handle_request(request_clone, state_clone).await;
let response_message = Message::new(response);
match response_message.encode() {
Ok(encoded) => {
if let Err(e) = MessageChannel::send(&*channel_clone, encoded).await {
error!("Failed to send HTTP-over-TCP response: {}", e);
}
}
Err(e) => {
error!("Failed to encode HTTP-over-TCP response: {}", e);
}
}
});
}
Err(e) => {
error!("Failed to receive HTTP-over-TCP request: {}", e);
}
}
}
})
};
Ok(server_handle)
}