use thiserror::Error;
#[cfg(feature = "ipc_channel")]
use crate::ipc::IpcManager;
use crate::ipc_types::{IpcHttpResponse, IpcPortNegotiation};
use axum::body::Body as AxumBody;
use axum::http::Request as AxumRequest;
use tower::util::ServiceExt;
use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::{Arc, Mutex as StdMutex};
use http_body_util::BodyExt;
use tracing::{info, debug, warn, error};
lazy_static::lazy_static! {
static ref PRE_ALLOCATED_PORT: Arc<StdMutex<Option<u16>>> = Arc::new(StdMutex::new(None));
}
pub fn set_pre_allocated_port(port: u16) {
if let Ok(mut guard) = PRE_ALLOCATED_PORT.lock() {
*guard = Some(port);
debug!("Set pre-allocated port: {}", port);
}
}
pub fn get_pre_allocated_port() -> Option<u16> {
PRE_ALLOCATED_PORT.lock().ok().and_then(|guard| *guard)
}
#[derive(Debug, Error)]
pub enum Error {
#[error("IO error: {0}")]
Io(#[from] std::io::Error),
#[error("Announcement error: {0}")]
Announce(String),
#[error("Configuration error: {0}")]
Config(String),
#[error("Server error: {0}")]
Server(String),
#[error("Internal error: {0}")]
Internal(Box<dyn std::error::Error + Send + Sync>),
}
#[cfg(feature = "ipc_channel")]
async fn serve_ipc(app: axum::Router) -> Result<(), Error> {
let manager = IpcManager::new();
let mut rx = manager.subscribe_http_requests();
let svc = app.clone();
while let Ok(ipc_req) = rx.recv().await {
let mut builder = AxumRequest::builder()
.method(ipc_req.method.as_str())
.uri(&ipc_req.uri);
for (k, v) in &ipc_req.headers {
builder = builder.header(k, v);
}
let req = builder
.body(AxumBody::from(ipc_req.body.unwrap_or_default()))
.map_err(|e| Error::Server(e.to_string()))?;
let resp = svc
.clone()
.oneshot(req)
.await
.map_err(|e| Error::Server(e.to_string()))?;
let status = resp.status().as_u16();
let mut headers = HashMap::new();
for (k, v) in resp.headers() {
headers.insert(k.to_string(), v.to_str().unwrap_or_default().to_string());
}
let full = resp.into_body()
.collect()
.await
.map_err(|e| Error::Server(e.to_string()))?;
let bytes = full.to_bytes();
let ipc_resp = IpcHttpResponse {
request_id: ipc_req.request_id.clone(),
status_code: status,
headers,
body: Some(bytes.to_vec()),
};
manager.send_http_response(ipc_resp)
.await
.map_err(Error::Server)?;
}
Ok(())
}
#[derive(Debug, Clone)]
pub struct ServeOptions {
pub bind_http: bool,
pub specific_port: Option<u16>,
pub listen_addr: Option<String>,
}
impl Default for ServeOptions {
fn default() -> Self {
Self {
bind_http: true,
specific_port: None,
listen_addr: None,
}
}
}
#[cfg(feature = "ipc_channel")]
async fn negotiate_port(specific_port: Option<u16>) -> Result<u16, Error> {
if let Some(pre_allocated) = get_pre_allocated_port() {
info!("Using pre-allocated port from InitBlob: {}", pre_allocated);
return Ok(pre_allocated);
}
if let Some(port) = specific_port {
warn!("No pre-allocated port found, attempting to negotiate specific port: {}", port);
} else {
warn!("No pre-allocated port found, attempting to negotiate dynamic port");
}
let manager = IpcManager::new();
let request = IpcPortNegotiation {
request_id: uuid::Uuid::new_v4().to_string(),
specific_port,
};
info!("Requesting port from orchestrator: {:?}", specific_port);
manager.send_port_negotiation(request)
.await
.map_err(|e| Error::Config(format!("Failed to send port request: {}", e)))?;
let response = manager.wait_for_port_response()
.await
.map_err(|e| Error::Config(format!("Failed to receive port response: {}", e)))?;
if !response.success {
return Err(Error::Config(response.error_message.unwrap_or_else(||
"Port negotiation failed with no error message".to_string())));
}
info!("Received port from orchestrator: {}", response.port);
Ok(response.port)
}
#[cfg(feature = "ipc_channel")]
pub async fn serve_with_options(app: axum::Router, options: ServeOptions) -> Result<(), Error> {
if !options.bind_http {
if std::env::var("IPC_ONLY").is_err() {
std::env::set_var("IPC_ONLY", "true");
info!("Set IPC_ONLY=true environment variable for SDK components");
}
if std::env::var("PYWATT_IPC_ONLY").is_err() {
std::env::set_var("PYWATT_IPC_ONLY", "true");
info!("Set PYWATT_IPC_ONLY=true environment variable for SDK components");
}
}
let ipc_task = tokio::spawn(serve_ipc(app.clone()));
if options.bind_http {
let port = negotiate_port(options.specific_port).await?;
let listen_addr = options.listen_addr.unwrap_or_else(|| "127.0.0.1".to_string());
let addr: SocketAddr = format!("{listen_addr}:{port}").parse()
.map_err(|e| Error::Config(format!("Invalid address or port: {}", e)))?;
info!("Starting HTTP server on {}", addr);
let listener = tokio::net::TcpListener::bind(&addr).await
.map_err(|e| Error::Server(e.to_string()))?;
let server = axum::serve(listener, app.into_make_service());
tokio::select! {
result = server => {
result.map_err(|e| Error::Server(e.to_string()))?;
}
result = ipc_task => {
result.map_err(|e| Error::Internal(Box::new(e)))??;
}
}
} else {
info!("Module serving via IPC only");
ipc_task.await.map_err(|e| Error::Internal(Box::new(e)))??;
}
Ok(())
}
#[cfg(not(feature = "ipc_channel"))]
pub async fn serve_with_options(app: axum::Router, options: ServeOptions) -> Result<(), Error> {
let listen_addr = options.listen_addr.unwrap_or_else(|| "127.0.0.1".to_string());
let port = options.specific_port.unwrap_or(0); let addr: SocketAddr = format!("{listen_addr}:{port}").parse()
.map_err(|e| Error::Config(format!("Invalid address or port: {}", e)))?;
info!("Starting HTTP server on {}", addr);
let listener = tokio::net::TcpListener::bind(&addr).await
.map_err(|e| Error::Server(e.to_string()))?;
let server = axum::serve(listener, app.into_make_service());
server.await.map_err(|e| Error::Server(e.to_string()))?;
Ok(())
}
#[cfg(feature = "ipc_channel")]
pub async fn serve_module(app: axum::Router) -> Result<(), Error> {
serve_with_options(app, ServeOptions::default()).await
}
#[cfg(not(feature = "ipc_channel"))]
pub async fn serve_module(app: axum::Router) -> Result<(), Error> {
serve_with_options(app, ServeOptions::default()).await
}
pub async fn serve_module_full<T, F, R>(
secret_keys: Vec<String>,
endpoints: Vec<crate::AnnouncedEndpoint>,
state_builder: F,
router_builder: R,
) -> Result<(), Error>
where
F: Fn(&crate::OrchestratorInit, Vec<secrecy::SecretString>) -> T + Send + Sync + 'static,
R: Fn(crate::AppState<T>) -> axum::Router + Send + Sync + 'static,
T: Send + Sync + Clone + 'static,
{
use crate::core::bootstrap::bootstrap_module;
let (app_state, ipc_handle) = bootstrap_module(
secret_keys,
endpoints,
state_builder,
None, ).await.map_err(|e| Error::Server(e.to_string()))?;
let router = router_builder(app_state);
let serve_task = tokio::spawn(serve_module(router));
tokio::select! {
result = serve_task => {
match result {
Ok(Ok(())) => {
info!("HTTP server completed successfully");
}
Ok(Err(e)) => {
error!("HTTP server error: {}", e);
return Err(e);
}
Err(e) => {
error!("HTTP server task panicked: {}", e);
return Err(Error::Internal(Box::new(e)));
}
}
}
result = ipc_handle => {
match result {
Ok(()) => {
info!("IPC processing completed (shutdown signal received)");
}
Err(e) => {
warn!("IPC processing task ended with error: {}", e);
}
}
}
}
info!("Module shutting down gracefully");
Ok(())
}
pub async fn serve_module_with_lifecycle<T, F, R>(
secret_keys: Vec<String>,
endpoints: Vec<crate::AnnouncedEndpoint>,
state_builder: F,
router_builder: R,
) -> Result<(), Error>
where
F: Fn(&crate::OrchestratorInit, Vec<secrecy::SecretString>) -> T + Send + Sync + 'static,
R: Fn(crate::AppState<T>) -> axum::Router + Send + Sync + 'static,
T: Send + Sync + Clone + 'static,
{
serve_module_full(secret_keys, endpoints, state_builder, router_builder).await
}