use std::sync::{Arc, OnceLock};
use async_trait::async_trait;
use authz_resolver_sdk::pep::PolicyEnforcer;
use toolkit::api::OpenApiRegistry;
use toolkit::{
DatabaseCapability, Gear, GearCtx, Healthcheck, HealthcheckResult, RestApiCapability,
};
use tracing::{debug, error, info, warn};
use graph_storage_sdk::GraphStorageClientV1;
use graph_storage_sdk::plugin_api::EmbeddingProviderV1;
use crate::api::rest::routes;
use crate::config::{EmbeddingProviderKind, GraphStorageConfig};
use crate::domain::embedding::SpaceState;
use crate::domain::local_client::GraphStorageLocalClient;
use crate::domain::service::GraphServices;
use crate::infra::embedding::fake::FakeEmbeddingProvider;
use crate::infra::engine::PgGraphEngine;
use crate::infra::store::{PgGraphStore, spaces};
#[toolkit::gear(name = "graph-storage", deps = [authz_resolver], capabilities = [db, rest])]
pub struct GraphStorage {
services: OnceLock<Arc<GraphServices>>,
}
impl Default for GraphStorage {
fn default() -> Self {
Self {
services: OnceLock::new(),
}
}
}
#[async_trait]
impl Gear for GraphStorage {
async fn init(&self, ctx: &GearCtx) -> anyhow::Result<()> {
let cfg = ctx.config_or_default::<GraphStorageConfig>()?.validated()?;
debug!(
traversal_hop = ?cfg.traversal_hop,
ingest_max_nodes = cfg.ingest_max_nodes,
"loaded graph-storage configuration"
);
let migrated = crate::infra::store::ingest::migrated_embedding_dimension();
if cfg.embedding_dimension != migrated {
anyhow::bail!(
"graph-storage.embedding_dimension is {} but the schema was migrated with {migrated}; \
vector search would compare incomparable vectors",
cfg.embedding_dimension
);
}
let db_raw = ctx.db_required()?;
let db = Arc::new(db_raw.db());
let pgq_available = crate::infra::engine::probe_pgq(&db).await;
if !pgq_available {
match cfg.traversal_hop {
crate::config::HopStrategy::Auto => warn!(
"this server does not provide SQL/PGQ; traversal will use the two-query hop"
),
crate::config::HopStrategy::Pgq => error!(
"traversal_hop is `pgq` and this server does not provide SQL/PGQ; the gear \
reports not ready and refuses traversal rather than substitute another \
backend"
),
crate::config::HopStrategy::TwoQuery => {}
}
}
let store = Arc::new(PgGraphStore::new(
Arc::clone(&db),
cfg.clone(),
pgq_available,
));
let engine = Arc::new(PgGraphEngine::new(Arc::clone(&store)));
let enforcer = PolicyEnforcer::new(ctx.client_hub().get()?);
let provider = select_embedding_provider(&cfg).await?;
let embedding = resolve_embedding_space(&db, provider, &cfg).await?;
let services = Arc::new(GraphServices::new(cfg, store, engine, enforcer, embedding));
self.services
.set(Arc::clone(&services))
.map_err(|_| anyhow::anyhow!("{} gear already initialized", Self::MODULE_NAME))?;
ctx.client_hub()
.register::<dyn GraphStorageClientV1>(Arc::new(GraphStorageLocalClient::new(services)));
info!(pgq_available, "graph-storage gear initialized");
Ok(())
}
}
async fn select_embedding_provider(
cfg: &GraphStorageConfig,
) -> anyhow::Result<Arc<dyn EmbeddingProviderV1>> {
let Some(kind) = cfg.embedding_provider else {
anyhow::bail!(
"graph-storage.embedding_provider is not set; name one of `fake`, `onnx` or \
`remote` -- the gear does not fall back to the fake, whose vectors rank nothing \
meaningfully"
);
};
match kind {
EmbeddingProviderKind::Fake => {
warn!(
"graph-storage.embedding_provider is `fake`: vector search will answer, \
but its ranking carries no semantics"
);
Ok(Arc::new(FakeEmbeddingProvider::new(
cfg.embedding_dimension,
)))
}
EmbeddingProviderKind::Onnx => onnx_provider(cfg).await,
EmbeddingProviderKind::Remote => remote_provider(cfg),
}
}
#[cfg(feature = "remote")]
fn remote_provider(cfg: &GraphStorageConfig) -> anyhow::Result<Arc<dyn EmbeddingProviderV1>> {
let named = |key: &str, value: &Option<String>| -> anyhow::Result<String> {
value.clone().ok_or_else(|| {
anyhow::anyhow!("graph-storage.{key} is required by the `remote` embedding provider")
})
};
let mut config = remote_embedding_plugin::RemoteProviderConfig::new(
named("embedding_remote_base_url", &cfg.embedding_remote_base_url)?,
named("embedding_remote_model", &cfg.embedding_remote_model)?,
);
config.dimension = cfg.embedding_dimension;
config.request_dimensions = cfg.embedding_remote_request_dimensions;
config.batch_size = cfg.embedding_remote_batch_size as usize;
config.timeout = std::time::Duration::from_secs(cfg.embedding_remote_timeout_secs);
if let Some(variable) = &cfg.embedding_remote_api_key_env {
let value = std::env::var(variable).map_err(|_| {
anyhow::anyhow!(
"graph-storage.embedding_remote_api_key_env names {variable}, which is not set \
in this process's environment"
)
})?;
if value.trim().is_empty() {
anyhow::bail!(
"graph-storage.embedding_remote_api_key_env names {variable}, which is empty"
);
}
config = config.with_api_key(value);
}
let provider = remote_embedding_plugin::RemoteEmbeddingProvider::new(config)?;
info!(
endpoint = %provider.endpoint(),
model = %provider.embedding_space().model_artifact,
"configured the remote embedding provider"
);
Ok(Arc::new(provider))
}
#[cfg(not(feature = "remote"))]
fn remote_provider(_cfg: &GraphStorageConfig) -> anyhow::Result<Arc<dyn EmbeddingProviderV1>> {
anyhow::bail!(
"graph-storage.embedding_provider is `remote` but this binary was built without the \
`remote` feature; rebuild with it or choose another provider"
)
}
#[cfg(feature = "onnx")]
async fn onnx_provider(cfg: &GraphStorageConfig) -> anyhow::Result<Arc<dyn EmbeddingProviderV1>> {
let named = |key: &str, value: &Option<String>| -> anyhow::Result<String> {
value.clone().ok_or_else(|| {
anyhow::anyhow!("graph-storage.{key} is required by the `onnx` embedding provider")
})
};
let mut config = onnx_embedding_plugin::OnnxProviderConfig::new(
named("embedding_model_path", &cfg.embedding_model_path)?,
named("embedding_tokenizer_path", &cfg.embedding_tokenizer_path)?,
);
config.dimension = cfg.embedding_dimension;
let provider = onnx_embedding_plugin::OnnxEmbeddingProvider::load(config).await?;
info!(
model = %provider.embedding_space().model_artifact,
"loaded the in-process ONNX embedding provider"
);
Ok(Arc::new(provider))
}
#[cfg(not(feature = "onnx"))]
#[expect(
clippy::unused_async,
reason = "one signature for both builds; the feature-enabled arm is async"
)]
async fn onnx_provider(_cfg: &GraphStorageConfig) -> anyhow::Result<Arc<dyn EmbeddingProviderV1>> {
anyhow::bail!(
"graph-storage.embedding_provider is `onnx` but this binary was built without the \
`onnx` feature; rebuild with it or choose another provider"
)
}
async fn resolve_embedding_space(
db: &toolkit_db::secure::Db,
provider: Arc<dyn EmbeddingProviderV1>,
cfg: &GraphStorageConfig,
) -> anyhow::Result<crate::domain::embedding::EmbeddingCoordinator> {
if provider.dimension() != cfg.embedding_dimension {
anyhow::bail!(
"the embedding provider declares {} dimensions but \
graph-storage.embedding_dimension is {}",
provider.dimension(),
cfg.embedding_dimension
);
}
let state = match spaces::resolve(db, provider.embedding_space()).await? {
spaces::SpaceResolution::Active { epoch } => {
info!(
epoch,
identity = %provider.embedding_space().identity_hash,
model = %provider.embedding_space().model_artifact,
"embedding space active"
);
SpaceState::Active { epoch }
}
spaces::SpaceResolution::Mismatched {
recorded_identity,
recorded_epoch,
} => {
error!(
recorded_epoch,
recorded_identity = %recorded_identity,
active_identity = %provider.embedding_space().identity_hash,
"stored vectors belong to a different embedding space than the configured \
provider; vector search is blocked until the graph is re-embedded"
);
SpaceState::Blocked
}
};
Ok(crate::domain::embedding::EmbeddingCoordinator::new(
provider,
state,
cfg.embedding_input_max_bytes,
))
}
impl DatabaseCapability for GraphStorage {
fn migrations(&self) -> Vec<Box<dyn sea_orm_migration::MigrationTrait>> {
use sea_orm_migration::MigratorTrait;
crate::infra::storage::migrations::Migrator::migrations()
}
}
impl RestApiCapability for GraphStorage {
fn register_rest(
&self,
_ctx: &GearCtx,
router: axum::Router,
openapi: &dyn OpenApiRegistry,
) -> anyhow::Result<axum::Router> {
let services = self
.services
.get()
.ok_or_else(|| anyhow::anyhow!("graph-storage services are not initialized"))?
.clone();
Ok(routes::register_routes(router, openapi, services))
}
fn healthcheck(&self, _ctx: &GearCtx) -> Option<Arc<dyn Healthcheck>> {
let services = self.services.get()?.clone();
Some(Arc::new(PlatformReadiness { services }))
}
}
struct PlatformReadiness {
services: Arc<GraphServices>,
}
#[async_trait]
impl Healthcheck for PlatformReadiness {
fn name(&self) -> &'static str {
"graph-storage"
}
async fn check(&self) -> HealthcheckResult {
platform_result(&self.services.readiness().await)
}
}
fn platform_result(readiness: &graph_storage_sdk::models::Readiness) -> HealthcheckResult {
use graph_storage_sdk::models::ReadinessState;
let named = |fatal: bool| {
readiness
.components
.iter()
.filter(|row| {
if fatal {
row.fatal()
} else {
matches!(
row.state,
ReadinessState::Degraded | ReadinessState::Unhealthy
)
}
})
.map(|row| row.component.as_str())
.collect::<Vec<_>>()
.join(", ")
};
if !readiness.ready {
return HealthcheckResult::unhealthy(format!("not ready: {}", named(true)))
.with_code("graph_storage.not_ready");
}
let degraded = named(false);
if degraded.is_empty() {
HealthcheckResult::healthy()
} else {
HealthcheckResult::degraded(format!("degraded: {degraded}"))
.with_code("graph_storage.degraded")
}
}
#[cfg(test)]
mod platform_readiness_tests {
use graph_storage_sdk::models::{
ComponentReadiness, DATABASE, DYNAMIC_INDEXES, EMBEDDING_SPACE, Readiness, ReadinessState,
};
use toolkit::HealthcheckStatus;
use super::platform_result;
fn row(component: &str, state: ReadinessState) -> ComponentReadiness {
ComponentReadiness::new(
component,
state,
"a problem text the platform must not see",
"what it blocks",
"recovery",
)
}
#[test]
fn all_healthy_is_healthy_and_a_missing_capability_is_not_a_fault() {
let result = platform_result(&Readiness::of(vec![
ComponentReadiness::healthy(DATABASE),
row(DYNAMIC_INDEXES, ReadinessState::NotImplemented),
]));
assert_eq!(result.status, HealthcheckStatus::Healthy, "{result:?}");
assert_eq!(result.code, None);
}
#[test]
fn a_space_mismatch_degrades_and_names_only_the_component() {
let result = platform_result(&Readiness::of(vec![
ComponentReadiness::healthy(DATABASE),
row(EMBEDDING_SPACE, ReadinessState::Unhealthy),
]));
assert_eq!(result.status, HealthcheckStatus::Degraded, "{result:?}");
assert_eq!(result.code.as_deref(), Some("graph_storage.degraded"));
let message = result.message.unwrap_or_default();
assert!(message.contains(EMBEDDING_SPACE), "{message}");
assert!(!message.contains("problem text"), "{message}");
}
#[test]
fn not_ready_is_unhealthy_and_names_what_blocks_it() {
let result = platform_result(&Readiness::of(vec![
row(DATABASE, ReadinessState::Unhealthy),
row(EMBEDDING_SPACE, ReadinessState::Unhealthy),
]));
assert_eq!(result.status, HealthcheckStatus::Unhealthy, "{result:?}");
assert_eq!(result.code.as_deref(), Some("graph_storage.not_ready"));
assert_eq!(result.message, Some(format!("not ready: {DATABASE}")));
}
}
#[cfg(test)]
mod tests {
use super::{EmbeddingProviderKind, GraphStorageConfig, select_embedding_provider};
fn with(kind: Option<EmbeddingProviderKind>) -> GraphStorageConfig {
GraphStorageConfig {
embedding_provider: kind,
..GraphStorageConfig::default()
}
}
#[tokio::test]
async fn an_unset_provider_is_refused_and_says_what_to_set() {
let refused = select_embedding_provider(&with(None))
.await
.err()
.expect("no provider is not a default");
let message = refused.to_string();
for named in ["embedding_provider", "fake", "onnx", "remote"] {
assert!(message.contains(named), "{named} is named: {message}");
}
}
#[tokio::test]
async fn the_fake_is_wired_at_the_configured_dimension() {
let config = GraphStorageConfig {
embedding_dimension: 16,
..with(Some(EmbeddingProviderKind::Fake))
};
let provider = select_embedding_provider(&config)
.await
.expect("the fake needs nothing");
assert_eq!(provider.embedding_space().dimension, 16);
}
#[cfg(feature = "remote")]
#[tokio::test]
async fn remote_is_wired_from_the_configuration() {
let configured = GraphStorageConfig {
embedding_dimension: 32,
embedding_remote_base_url: Some("http://127.0.0.1:9/v1".to_owned()),
embedding_remote_model: Some("wired-model".to_owned()),
..with(Some(EmbeddingProviderKind::Remote))
};
let provider = select_embedding_provider(&configured)
.await
.expect("a complete remote configuration boots");
assert_eq!(
provider.embedding_space().model_artifact,
"wired-model@http://127.0.0.1:9/v1/embeddings"
);
assert_eq!(provider.embedding_space().dimension, 32);
for (missing, key) in [
(
GraphStorageConfig {
embedding_remote_base_url: None,
..configured.clone()
},
"embedding_remote_base_url",
),
(
GraphStorageConfig {
embedding_remote_model: None,
..configured.clone()
},
"embedding_remote_model",
),
] {
let refused = select_embedding_provider(&missing)
.await
.err()
.expect("a required key is required");
assert!(refused.to_string().contains(key), "{key}: {refused}");
}
let unset = GraphStorageConfig {
embedding_remote_api_key_env: Some(
"GRAPH_STORAGE_TEST_CREDENTIAL_THAT_IS_NEVER_SET".to_owned(),
),
..configured
};
let refused = select_embedding_provider(&unset)
.await
.err()
.expect("an unset credential variable stops the boot");
assert!(
refused.to_string().contains("not set"),
"the refusal says why: {refused}"
);
}
#[cfg(not(feature = "remote"))]
#[tokio::test]
async fn remote_without_the_feature_says_so() {
let refused = select_embedding_provider(&with(Some(EmbeddingProviderKind::Remote)))
.await
.err()
.expect("a provider the binary lacks is refused");
assert!(
refused.to_string().contains("`remote` feature"),
"{refused}"
);
}
#[cfg(feature = "onnx")]
#[tokio::test]
async fn onnx_names_the_artifact_it_is_missing() {
let refused = select_embedding_provider(&with(Some(EmbeddingProviderKind::Onnx)))
.await
.err()
.expect("no artifacts, no provider");
assert!(
refused.to_string().contains("embedding_model_path"),
"{refused}"
);
}
#[cfg(feature = "onnx")]
#[tokio::test]
async fn onnx_is_wired_from_the_configuration() {
let (Ok(model), Ok(tokenizer)) = (
std::env::var("GRAPH_STORAGE_ONNX_MODEL"),
std::env::var("GRAPH_STORAGE_ONNX_TOKENIZER"),
) else {
assert!(
std::env::var("GRAPH_STORAGE_ONNX_REQUIRED").is_err(),
"GRAPH_STORAGE_ONNX_REQUIRED is set but the model artifacts are not"
);
eprintln!("no ONNX artifacts in this environment - skipping");
return;
};
let configured = GraphStorageConfig {
embedding_model_path: Some(model),
embedding_tokenizer_path: Some(tokenizer),
..with(Some(EmbeddingProviderKind::Onnx))
};
let provider = select_embedding_provider(&configured)
.await
.expect("the downloaded artifacts load");
assert_eq!(
provider.embedding_space().dimension,
configured.embedding_dimension
);
}
#[cfg(not(feature = "onnx"))]
#[tokio::test]
async fn onnx_without_the_feature_says_so() {
let refused = select_embedding_provider(&with(Some(EmbeddingProviderKind::Onnx)))
.await
.err()
.expect("a provider the binary lacks is refused");
assert!(refused.to_string().contains("`onnx` feature"), "{refused}");
}
}