river-data-core 0.4.0

Common types & traits for the in the river-data platform
Documentation
use axum::extract::State;
use axum::Json;
use chrono::Utc;
use moka::future::Cache;
use sea_orm::{ActiveModelTrait, ColumnTrait, Condition, ConnectionTrait, EntityTrait, QueryFilter, Set, Statement, DatabaseBackend};
use std::sync::LazyLock;
use std::time::Duration;
use uuid::Uuid;

use super::{SyncError, SyncResult};
use crate::models::{CommandStatus, HeartbeatRequest, HeartbeatResponse, PendingCommand, ServiceStatus};
use crate::server::entity::{sync_commands, sync_services};
use crate::server::handlers::enroll::create_session_token;
use crate::server::middleware::SyncServiceContext;
use crate::server::state::SyncState;

pub(crate) static SESSION_TOKEN_CACHE: LazyLock<Cache<Uuid, String>> = LazyLock::new(|| {
    Cache::builder()
        .max_capacity(100)
        .time_to_live(Duration::from_secs(13 * 60))
        .build()
});

/// Periodic heartbeat from a sync service. Updates `last_heartbeat`, `status`, and
/// `current_operation`. Returns a fresh session token (rotated on every heartbeat) and
/// any pending commands queued for this service. Requires sync session token auth.
#[utoipa::path(
    post,
    path = "/heartbeat",
    request_body = HeartbeatRequest,
    responses(
        (status = 200, description = "Heartbeat acknowledged; fresh token and pending commands", body = HeartbeatResponse),
        (status = 400, description = "Invalid status string"),
        (status = 401, description = "Invalid or expired session token"),
    ),
    tag = "sync"
)]
pub async fn heartbeat<S: SyncState>(
    State(state): State<S>,
    _ctx: SyncServiceContext,
    Json(req): Json<HeartbeatRequest>,
) -> SyncResult<Json<HeartbeatResponse>> {
    if ServiceStatus::from_str(&req.status).is_none() {
        let valid: Vec<&str> = ServiceStatus::ALL.iter().map(|s| s.as_str()).collect();
        return Err(SyncError::BadRequest(format!(
            "Invalid status '{}'. Valid: {}",
            req.status,
            valid.join(", ")
        )));
    }

    let service = sync_services::Entity::find_by_id(req.service_id)
        .one(state.db())
        .await?
        .ok_or_else(|| SyncError::NotFound("Service not found".to_string()))?;

    let mut active: sync_services::ActiveModel = service.into();
    active.status = Set(req.status);
    active.current_operation = Set(req.current_operation);
    active.last_heartbeat = Set(Some(Utc::now().into()));
    active.updated_at = Set(Utc::now().into());
    active.update(state.db()).await?;

    let session_token = if let Some(cached) = SESSION_TOKEN_CACHE.get(&req.service_id).await {
        cached
    } else {
        let token = create_session_token(&state, req.service_id).await?;
        SESSION_TOKEN_CACHE
            .insert(req.service_id, token.clone())
            .await;
        token
    };

    let pending = sync_commands::Entity::find()
        .filter(
            Condition::all()
                .add(sync_commands::Column::ServiceId.eq(req.service_id))
                .add(sync_commands::Column::Status.eq(CommandStatus::Pending.as_str()))
                .add(sync_commands::Column::ExpiresAt.gt(Utc::now())),
        )
        .all(state.db())
        .await?;

    let pending_commands = pending
        .into_iter()
        .map(|c| PendingCommand {
            id: c.id,
            command: c.command,
            payload: c.payload,
        })
        .collect();

    let db_clone = state.db().clone();
    let sid = req.service_id;
    tokio::spawn(async move {
        let _ = db_clone
            .execute(Statement::from_sql_and_values(
                DatabaseBackend::Postgres,
                "UPDATE sync_commands SET status = 'expired' WHERE service_id = $1 AND status = 'pending' AND expires_at < NOW()",
                [sid.into()],
            ))
            .await;
    });

    Ok(Json(HeartbeatResponse {
        session_token,
        pending_commands,
    }))
}