use std::future::Future;
use std::net::SocketAddr;
use std::sync::Arc;
use arrow_flight::flight_service_server::FlightServiceServer;
use async_trait::async_trait;
use axum::routing::get;
use axum::Router;
use datafusion::execution::context::SessionContext;
use datafusion_flight_sql_server::service::FlightSqlService;
use jammi_ai::session::InferenceSession;
use jammi_db::config::JammiConfig;
use tokio::net::TcpListener;
use tokio::signal;
use tokio::sync::broadcast;
use tonic::transport::Server;
use tonic_web::GrpcWebLayer;
use crate::error::fallback_handler;
use crate::flight::TenantBoundProvider;
use crate::grpc::audit::AuditServer;
use crate::grpc::catalog::CatalogServer;
use crate::grpc::embedding::EmbeddingServer;
use crate::grpc::eval::EvalServer;
use crate::grpc::inference::InferenceServer;
use crate::grpc::pipeline::PipelineServer;
use crate::grpc::proto::audit::audit_service_server::AuditServiceServer;
use crate::grpc::proto::catalog::catalog_service_server::CatalogServiceServer;
use crate::grpc::proto::embedding::embedding_service_server::EmbeddingServiceServer;
use crate::grpc::proto::eval::eval_service_server::EvalServiceServer;
use crate::grpc::proto::inference::inference_service_server::InferenceServiceServer;
use crate::grpc::proto::pipeline::pipeline_service_server::PipelineServiceServer;
#[cfg(feature = "train")]
use crate::grpc::proto::training::training_service_server::TrainingServiceServer;
use crate::grpc::proto::trigger::trigger_service_server::TriggerServiceServer;
use crate::grpc::session::{SessionStore, TenantInterceptor};
#[cfg(feature = "train")]
use crate::grpc::training::TrainingServer;
use crate::grpc::trigger::TriggerServer;
use crate::grpc_web_trailers::GrpcWebTrailersLayer;
use crate::metrics_layer::MetricsLayer;
use crate::routes::health::{self, MetricsRegistry};
use crate::tiers::{ServiceTier, TierSet};
#[derive(Debug, thiserror::Error)]
pub enum ServerError {
#[error("config error: {0}")]
Config(String),
#[error("service tier: {0}")]
Tier(#[from] crate::tiers::TierError),
#[error("engine init: {0}")]
Engine(#[from] jammi_db::error::JammiError),
#[error("metrics registry: {0}")]
Metrics(#[from] prometheus::Error),
#[error("transport: {0}")]
Transport(#[from] tonic::transport::Error),
#[error("io: {0}")]
Io(#[from] std::io::Error),
#[error("addr parse: {0}")]
AddrParse(#[from] std::net::AddrParseError),
}
#[async_trait]
pub trait ReadinessCheck: Send + Sync {
async fn check(&self) -> Result<(), String>;
}
pub struct ReadinessProbe {
inner: Arc<dyn ReadinessCheck>,
}
impl ReadinessProbe {
pub fn new(inner: Arc<dyn ReadinessCheck>) -> Self {
Self { inner }
}
pub async fn check(&self) -> Result<(), String> {
self.inner.check().await
}
}
pub struct CatalogPingProbe {
session: Arc<InferenceSession>,
}
impl CatalogPingProbe {
pub fn new(session: Arc<InferenceSession>) -> Self {
Self { session }
}
}
#[async_trait]
impl ReadinessCheck for CatalogPingProbe {
async fn check(&self) -> Result<(), String> {
self.session
.catalog()
.ping()
.await
.map_err(|e| e.to_string())
}
}
pub struct OssServer {
flight_addr: SocketAddr,
health_addr: SocketAddr,
session: Arc<InferenceSession>,
session_store: SessionStore,
metrics: Arc<MetricsRegistry>,
readiness: Arc<ReadinessProbe>,
tiers: TierSet,
}
impl OssServer {
pub async fn new(config: JammiConfig) -> Result<Self, ServerError> {
config
.server
.validate()
.map_err(|e| ServerError::Config(e.to_string()))?;
config
.training
.worker_intervals()
.map_err(|e| ServerError::Config(e.to_string()))?;
let flight_addr: SocketAddr = config.server.flight_listen.parse()?;
let health_addr: SocketAddr = config.server.health_listen.parse()?;
let tiers = TierSet::from_config(&config.server.services)?;
let session = InferenceSession::open(config).await?;
let session_store = SessionStore::new();
let metrics = Arc::new(MetricsRegistry::new()?);
let readiness = Arc::new(ReadinessProbe::new(Arc::new(CatalogPingProbe::new(
Arc::clone(&session),
))));
Ok(Self {
flight_addr,
health_addr,
session,
session_store,
metrics,
readiness,
tiers,
})
}
pub fn flight_addr(&self) -> SocketAddr {
self.flight_addr
}
pub fn health_addr(&self) -> SocketAddr {
self.health_addr
}
pub fn metrics(&self) -> Arc<MetricsRegistry> {
Arc::clone(&self.metrics)
}
pub fn session(&self) -> Arc<InferenceSession> {
Arc::clone(&self.session)
}
pub fn with_readiness(mut self, readiness: Arc<ReadinessProbe>) -> Self {
self.readiness = readiness;
self
}
pub async fn run(self) -> Result<(), ServerError> {
self.run_with_shutdown(shutdown_signal()).await
}
pub async fn run_with_shutdown(
self,
shutdown: impl Future<Output = ()> + Send + 'static,
) -> Result<(), ServerError> {
let (shutdown_tx, _) = broadcast::channel::<()>(1);
let mut shutdown_health_rx = shutdown_tx.subscribe();
let mut shutdown_grpc_rx = shutdown_tx.subscribe();
let shutdown_tx_for_signal = shutdown_tx.clone();
tokio::spawn(async move {
shutdown.await;
let _ = shutdown_tx_for_signal.send(());
});
let health_router = self.build_health_router();
let health_listener = TcpListener::bind(self.health_addr).await?;
tracing::info!(
address = %self.health_addr,
"HTTP side-channel listening (/healthz, /readyz, /metrics)"
);
let health_task = tokio::spawn(async move {
axum::serve(health_listener, health_router)
.with_graceful_shutdown(async move {
let _ = shutdown_health_rx.recv().await;
})
.await
.map_err(ServerError::from)
});
let grpc_future = self.build_and_serve_grpc(async move {
let _ = shutdown_grpc_rx.recv().await;
});
let grpc_result = grpc_future.await;
if grpc_result.is_err() {
let _ = shutdown_tx.send(());
}
let health_result = match health_task.await {
Ok(r) => r,
Err(join_err) => Err(ServerError::Io(std::io::Error::other(join_err.to_string()))),
};
grpc_result.and(health_result)
}
fn build_health_router(&self) -> Router {
let readyz = Router::new()
.route("/readyz", get(health::readyz))
.with_state(Arc::clone(&self.readiness));
let metrics = Router::new()
.route("/metrics", get(health::metrics))
.with_state(Arc::clone(&self.metrics));
Router::new()
.route("/healthz", get(health::healthz))
.merge(readyz)
.merge(metrics)
.fallback(fallback_handler)
}
async fn build_and_serve_grpc(
&self,
shutdown: impl Future<Output = ()> + Send + 'static,
) -> Result<(), ServerError> {
let trigger = self
.tiers
.contains(ServiceTier::Event)
.then(|| crate::TriggerHandles {
topic_repo: self.session.topic_repo(),
publisher: self.session.publisher(),
subscriber: self.session.subscriber(),
});
serve_grpc_chain(
GrpcChain {
addr: self.flight_addr,
flight_ctx: self.session.context().clone(),
flight_binding: self.session.tenant_binding_arc(),
store: self.session_store.clone(),
trigger,
engine: Some(Arc::clone(&self.session)),
tiers: self.tiers.clone(),
metrics: Arc::clone(&self.metrics),
},
shutdown,
)
.await
}
}
pub struct GrpcChain {
pub addr: SocketAddr,
pub flight_ctx: SessionContext,
pub flight_binding: jammi_db::tenant_scope::TenantBinding,
pub store: SessionStore,
pub trigger: Option<crate::TriggerHandles>,
pub engine: Option<Arc<InferenceSession>>,
pub tiers: TierSet,
pub metrics: Arc<MetricsRegistry>,
}
pub async fn serve_grpc_chain(
chain: GrpcChain,
shutdown: impl Future<Output = ()> + Send + 'static,
) -> Result<(), ServerError> {
let GrpcChain {
addr,
flight_ctx,
flight_binding,
store,
trigger,
engine,
tiers,
metrics,
} = chain;
let interceptor = TenantInterceptor::new(store.clone());
let provider = TenantBoundProvider::new(flight_ctx.state(), flight_binding, store.clone());
let flight = FlightSqlService::new_with_provider(Box::new(provider));
let flight_svc = FlightServiceServer::new(flight);
let catalog_svc = CatalogServiceServer::with_interceptor(
CatalogServer::new(store, tiers.clone(), engine.clone()),
interceptor.clone(),
);
let mut builder = Server::builder()
.accept_http1(true)
.layer(MetricsLayer::new(metrics))
.layer(GrpcWebTrailersLayer::new())
.layer(GrpcWebLayer::new())
.add_service(flight_svc)
.add_service(catalog_svc);
let mut mounted = vec!["Flight SQL", "CatalogService"];
if let Some(handles) = trigger {
let trigger_svc = TriggerServiceServer::with_interceptor(
TriggerServer::new(handles.topic_repo, handles.publisher, handles.subscriber),
interceptor.clone(),
);
builder = builder.add_service(trigger_svc);
mounted.push("TriggerService");
}
#[cfg(feature = "train")]
let mut _train_worker: Option<jammi_ai::fine_tune::worker::EmbeddedWorker> = None;
if let Some(session) = engine {
let embedding_svc = EmbeddingServiceServer::with_interceptor(
EmbeddingServer::new(Arc::clone(&session)),
interceptor.clone(),
);
builder = builder.add_service(embedding_svc);
mounted.push("EmbeddingService");
let inference_svc = InferenceServiceServer::with_interceptor(
InferenceServer::new(Arc::clone(&session)),
interceptor.clone(),
);
builder = builder.add_service(inference_svc);
mounted.push("InferenceService");
let pipeline_svc = PipelineServiceServer::with_interceptor(
PipelineServer::new(Arc::clone(&session)),
interceptor.clone(),
);
builder = builder.add_service(pipeline_svc);
mounted.push("PipelineService");
let audit_svc = AuditServiceServer::with_interceptor(
AuditServer::new(Arc::clone(&session)),
interceptor.clone(),
);
builder = builder.add_service(audit_svc);
mounted.push("AuditService");
if tiers.contains(ServiceTier::Eval) {
let eval_svc = EvalServiceServer::with_interceptor(
EvalServer::new(Arc::clone(&session)),
interceptor.clone(),
);
builder = builder.add_service(eval_svc);
mounted.push("EvalService");
}
#[cfg(feature = "train")]
if tiers.contains(ServiceTier::Train) {
_train_worker = Some(jammi_ai::fine_tune::worker::EmbeddedWorker::spawn(
&session,
)?);
let training_svc =
TrainingServiceServer::with_interceptor(TrainingServer::new(session), interceptor);
builder = builder.add_service(training_svc);
mounted.push("TrainingService");
}
}
tracing::info!("gRPC chain ({}) listening on {addr}", mounted.join(" + "));
builder
.serve_with_shutdown(addr, shutdown)
.await
.map_err(ServerError::from)
}
async fn shutdown_signal() {
let ctrl_c = async {
match signal::ctrl_c().await {
Ok(()) => {}
Err(e) => tracing::error!("Failed to install Ctrl+C handler: {e}"),
}
};
#[cfg(unix)]
let terminate = async {
match signal::unix::signal(signal::unix::SignalKind::terminate()) {
Ok(mut sig) => {
sig.recv().await;
}
Err(e) => {
tracing::error!("Failed to install SIGTERM handler: {e}");
std::future::pending::<()>().await;
}
}
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! {
() = ctrl_c => {},
() = terminate => {},
}
tracing::info!("Shutdown signal received, draining connections...");
}