use std::sync::Arc;
use axum::extract::FromRef;
use sqlx::PgPool;
use tokio::sync::Semaphore;
use tokio_util::task::TaskTracker;
use uuid::Uuid;
use yorishiro_core::ResultExt;
use yorishiro_core::db::TenantDb;
use yorishiro_core::repositories::entities::EntityRecord;
use yorishiro_core::services::embedding::EmbeddingProvider;
use yorishiro_core::services::embedding::sync as embedding_sync;
const EMBEDDING_SYNC_MAX_CONCURRENCY: usize = 4;
#[derive(Clone)]
pub struct AppState {
pub tenant_db: TenantDb,
pub identity_pool: PgPool,
pub embedding_provider: Arc<dyn EmbeddingProvider>,
embedding_sync_permits: Arc<Semaphore>,
embedding_tasks: TaskTracker,
}
impl AppState {
pub fn new(
tenant_db: TenantDb,
identity_pool: PgPool,
embedding_provider: Arc<dyn EmbeddingProvider>,
) -> Self {
Self {
tenant_db,
identity_pool,
embedding_provider,
embedding_sync_permits: Arc::new(Semaphore::new(EMBEDDING_SYNC_MAX_CONCURRENCY)),
embedding_tasks: TaskTracker::new(),
}
}
pub fn embedding_tasks(&self) -> &TaskTracker {
&self.embedding_tasks
}
pub fn spawn_embedding_sync(
&self,
tenant_id: Uuid,
workspace_id: Uuid,
record: EntityRecord,
) -> tokio::task::JoinHandle<()> {
let db = self.tenant_db.clone();
let provider = Arc::clone(&self.embedding_provider);
let permits = Arc::clone(&self.embedding_sync_permits);
self.embedding_tasks.spawn(async move {
let Ok(_permit) = permits.acquire_owned().await else {
return;
};
let result = async {
let mut conn = db
.acquire_for_workspace(tenant_id, workspace_id)
.await
.internal()?;
embedding_sync::sync_embedding_for_record(
&mut conn,
workspace_id,
&record,
provider.as_ref(),
)
.await
}
.await;
if let Err(err) = result {
tracing::warn!(entity_id = %record.id, error = %err, "embedding sync failed");
}
})
}
}
impl FromRef<AppState> for TenantDb {
fn from_ref(state: &AppState) -> Self {
state.tenant_db.clone()
}
}