use crate::error::{ErrorData, Result};
use alien_core::{CommandDeliveryMode, CommandState, CommandTarget, CommandTargetType};
use alien_error::AlienError;
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use uuid::Uuid;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ResolvedCommandTarget {
pub target: CommandTarget,
pub delivery_mode: CommandDeliveryMode,
}
#[derive(Debug, Clone)]
pub struct CommandMetadata {
pub command_id: String,
pub target: CommandTarget,
pub delivery_mode: CommandDeliveryMode,
pub project_id: String,
}
#[derive(Debug, Clone)]
pub struct CommandEnvelopeData {
pub command_id: String,
pub deployment_id: String,
pub command: String, pub attempt: u32,
pub deadline: Option<DateTime<Utc>>,
pub state: CommandState,
pub target: CommandTarget,
pub delivery_mode: CommandDeliveryMode,
}
#[derive(Debug, Clone)]
pub struct CommandStatus {
pub command_id: String,
pub deployment_id: String,
pub command: String, pub state: CommandState,
pub attempt: u32,
pub deadline: Option<DateTime<Utc>>,
pub created_at: DateTime<Utc>,
pub dispatched_at: Option<DateTime<Utc>>,
pub completed_at: Option<DateTime<Utc>>,
pub error: Option<serde_json::Value>,
pub request_size_bytes: Option<u64>,
pub response_size_bytes: Option<u64>,
pub target: CommandTarget,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
struct CommandRecord {
id: String,
deployment_id: String,
command: String,
state: CommandState,
attempt: u32,
deadline: Option<DateTime<Utc>>,
created_at: DateTime<Utc>,
dispatched_at: Option<DateTime<Utc>>,
completed_at: Option<DateTime<Utc>>,
request_size_bytes: Option<u64>,
response_size_bytes: Option<u64>,
error: Option<serde_json::Value>,
target: CommandTarget,
delivery_mode: CommandDeliveryMode,
project_id: String,
}
pub fn validate_command_target_id(resource_id: &str) -> Result<()> {
if resource_id.contains(':') {
return Err(AlienError::new(ErrorData::CommandTargetIdInvalid {
resource_id: resource_id.to_string(),
}));
}
Ok(())
}
pub fn validate_command_name(command: &str) -> Result<()> {
if command.contains(':') {
return Err(AlienError::new(ErrorData::InvalidCommand {
message: format!("Command name '{command}' must not contain ':'"),
}));
}
Ok(())
}
pub fn select_command_target(
deployment_id: &str,
targets: &[CommandTarget],
requested: Option<&str>,
) -> Result<CommandTarget> {
let target = match requested {
Some(resource_id) => {
validate_command_target_id(resource_id)?;
let found = if resource_id.is_empty() {
None
} else {
targets.iter().find(|t| t.resource_id == resource_id)
};
found
.ok_or_else(|| {
AlienError::new(ErrorData::CommandTargetNotFound {
resource_id: resource_id.to_string(),
deployment_id: deployment_id.to_string(),
})
})?
.clone()
}
None => match targets {
[] => {
return Err(AlienError::new(ErrorData::NoCommandTargets {
deployment_id: deployment_id.to_string(),
}))
}
[single] => single.clone(),
_ => {
return Err(AlienError::new(ErrorData::CommandTargetAmbiguous {
deployment_id: deployment_id.to_string(),
}))
}
},
};
validate_command_target_id(&target.resource_id)?;
Ok(target)
}
pub fn delivery_mode_for(
resource_type: CommandTargetType,
worker_mode: CommandDeliveryMode,
) -> CommandDeliveryMode {
match resource_type {
CommandTargetType::Container | CommandTargetType::Daemon => CommandDeliveryMode::Pull,
CommandTargetType::Worker => worker_mode,
}
}
#[async_trait]
pub trait CommandRegistry: Send + Sync {
async fn resolve_target(
&self,
deployment_id: &str,
requested: Option<&str>,
) -> Result<ResolvedCommandTarget>;
async fn create_command(
&self,
deployment_id: &str,
command_name: &str,
target: &ResolvedCommandTarget,
initial_state: CommandState,
deadline: Option<DateTime<Utc>>,
request_size_bytes: Option<u64>,
) -> Result<CommandMetadata>;
async fn get_command_metadata(&self, command_id: &str) -> Result<Option<CommandEnvelopeData>>;
async fn get_command_status(&self, command_id: &str) -> Result<Option<CommandStatus>>;
async fn update_command_state(
&self,
command_id: &str,
state: CommandState,
dispatched_at: Option<DateTime<Utc>>,
completed_at: Option<DateTime<Utc>>,
response_size_bytes: Option<u64>,
error: Option<serde_json::Value>,
) -> Result<bool>;
async fn complete_command(
&self,
command_id: &str,
state: CommandState,
completed_at: DateTime<Utc>,
response_size_bytes: Option<u64>,
error: Option<serde_json::Value>,
) -> Result<bool>;
async fn mark_dispatched_if_not_terminal(
&self,
command_id: &str,
dispatched_at: DateTime<Utc>,
) -> Result<bool>;
async fn increment_attempt(&self, command_id: &str) -> Result<u32>;
}
pub struct InMemoryCommandRegistry {
commands: Arc<RwLock<HashMap<String, CommandRecord>>>,
targets: Arc<RwLock<Vec<CommandTarget>>>,
worker_delivery_mode: CommandDeliveryMode,
}
impl InMemoryCommandRegistry {
pub fn new() -> Self {
Self::with_worker_delivery_mode(CommandDeliveryMode::Pull)
}
pub fn with_worker_delivery_mode(worker_delivery_mode: CommandDeliveryMode) -> Self {
Self {
commands: Arc::new(RwLock::new(HashMap::new())),
targets: Arc::new(RwLock::new(Vec::new())),
worker_delivery_mode,
}
}
pub async fn register_target(
&self,
resource_id: impl Into<String>,
resource_type: CommandTargetType,
) -> Result<()> {
let resource_id = resource_id.into();
validate_command_target_id(&resource_id)?;
self.targets
.write()
.await
.push(CommandTarget::new(resource_id, resource_type));
Ok(())
}
#[allow(dead_code)]
pub async fn list_command_ids(&self) -> Vec<String> {
let commands = self.commands.read().await;
commands.keys().cloned().collect()
}
}
impl Default for InMemoryCommandRegistry {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl CommandRegistry for InMemoryCommandRegistry {
async fn resolve_target(
&self,
deployment_id: &str,
requested: Option<&str>,
) -> Result<ResolvedCommandTarget> {
let targets = self.targets.read().await;
let target = select_command_target(deployment_id, &targets, requested)?;
let delivery_mode = delivery_mode_for(target.resource_type, self.worker_delivery_mode);
Ok(ResolvedCommandTarget {
target,
delivery_mode,
})
}
async fn create_command(
&self,
deployment_id: &str,
command_name: &str,
target: &ResolvedCommandTarget,
initial_state: CommandState,
deadline: Option<DateTime<Utc>>,
request_size_bytes: Option<u64>,
) -> Result<CommandMetadata> {
let command_id = format!("cmd_{}", Uuid::new_v4());
let record = CommandRecord {
id: command_id.clone(),
deployment_id: deployment_id.to_string(),
command: command_name.to_string(),
state: initial_state,
attempt: 1,
deadline,
created_at: Utc::now(),
dispatched_at: None,
completed_at: None,
request_size_bytes,
response_size_bytes: None,
error: None,
target: target.target.clone(),
delivery_mode: target.delivery_mode,
project_id: "local-dev".to_string(),
};
self.commands
.write()
.await
.insert(command_id.clone(), record);
Ok(CommandMetadata {
command_id,
target: target.target.clone(),
delivery_mode: target.delivery_mode,
project_id: "local-dev".to_string(),
})
}
async fn get_command_metadata(&self, command_id: &str) -> Result<Option<CommandEnvelopeData>> {
let commands = self.commands.read().await;
Ok(commands.get(command_id).map(|r| CommandEnvelopeData {
command_id: r.id.clone(),
deployment_id: r.deployment_id.clone(),
command: r.command.clone(),
attempt: r.attempt,
deadline: r.deadline,
state: r.state,
target: r.target.clone(),
delivery_mode: r.delivery_mode,
}))
}
async fn get_command_status(&self, command_id: &str) -> Result<Option<CommandStatus>> {
let commands = self.commands.read().await;
Ok(commands.get(command_id).map(|r| CommandStatus {
command_id: r.id.clone(),
deployment_id: r.deployment_id.clone(),
command: r.command.clone(),
state: r.state,
attempt: r.attempt,
deadline: r.deadline,
created_at: r.created_at,
dispatched_at: r.dispatched_at,
completed_at: r.completed_at,
error: r.error.clone(),
request_size_bytes: r.request_size_bytes,
response_size_bytes: r.response_size_bytes,
target: r.target.clone(),
}))
}
async fn update_command_state(
&self,
command_id: &str,
state: CommandState,
dispatched_at: Option<DateTime<Utc>>,
completed_at: Option<DateTime<Utc>>,
response_size_bytes: Option<u64>,
error: Option<serde_json::Value>,
) -> Result<bool> {
let mut commands = self.commands.write().await;
if let Some(record) = commands.get_mut(command_id) {
if record.state.is_terminal() {
return Ok(false);
}
record.state = state;
if let Some(ts) = dispatched_at {
record.dispatched_at = Some(ts);
}
if let Some(ts) = completed_at {
record.completed_at = Some(ts);
}
if let Some(size) = response_size_bytes {
record.response_size_bytes = Some(size);
}
if let Some(err) = error {
record.error = Some(err);
}
return Ok(true);
}
Ok(false)
}
async fn complete_command(
&self,
command_id: &str,
state: CommandState,
completed_at: DateTime<Utc>,
response_size_bytes: Option<u64>,
error: Option<serde_json::Value>,
) -> Result<bool> {
let mut commands = self.commands.write().await;
let Some(record) = commands.get_mut(command_id) else {
return Ok(false);
};
if record.state.is_terminal() {
return Ok(false);
}
record.state = state;
record.completed_at = Some(completed_at);
if let Some(size) = response_size_bytes {
record.response_size_bytes = Some(size);
}
if let Some(err) = error {
record.error = Some(err);
}
Ok(true)
}
async fn mark_dispatched_if_not_terminal(
&self,
command_id: &str,
dispatched_at: DateTime<Utc>,
) -> Result<bool> {
let mut commands = self.commands.write().await;
let Some(record) = commands.get_mut(command_id) else {
return Ok(false);
};
if record.state.is_terminal() {
return Ok(false);
}
record.state = CommandState::Dispatched;
record.dispatched_at = Some(dispatched_at);
Ok(true)
}
async fn increment_attempt(&self, command_id: &str) -> Result<u32> {
let mut commands = self.commands.write().await;
if let Some(record) = commands.get_mut(command_id) {
record.attempt += 1;
Ok(record.attempt)
} else {
Ok(1) }
}
}
#[cfg(test)]
mod tests {
use super::*;
use alien_core::{CommandDeliveryMode, CommandTarget, CommandTargetType};
async fn resolved(
registry: &InMemoryCommandRegistry,
requested: Option<&str>,
) -> Result<ResolvedCommandTarget> {
registry.resolve_target("dep-1", requested).await
}
#[tokio::test]
async fn test_resolve_explicit_target_found() {
let registry = InMemoryCommandRegistry::new();
registry
.register_target("worker-a", CommandTargetType::Worker)
.await
.unwrap();
registry
.register_target("daemon-b", CommandTargetType::Daemon)
.await
.unwrap();
let result = resolved(®istry, Some("daemon-b")).await.unwrap();
assert_eq!(
result.target,
CommandTarget::new("daemon-b", CommandTargetType::Daemon)
);
}
#[tokio::test]
async fn test_resolve_explicit_unknown_target_is_not_found() {
let registry = InMemoryCommandRegistry::new();
registry
.register_target("worker-a", CommandTargetType::Worker)
.await
.unwrap();
let err = resolved(®istry, Some("no-such-resource"))
.await
.unwrap_err();
assert_eq!(err.code, "COMMAND_TARGET_NOT_FOUND");
assert_eq!(err.http_status_code, Some(404));
}
#[tokio::test]
async fn test_resolve_explicit_empty_string_is_not_found() {
let registry = InMemoryCommandRegistry::new();
registry
.register_target("worker-a", CommandTargetType::Worker)
.await
.unwrap();
let err = resolved(®istry, Some("")).await.unwrap_err();
assert_eq!(err.code, "COMMAND_TARGET_NOT_FOUND");
}
#[tokio::test]
async fn test_resolve_shorthand_single_target_resolves() {
let registry = InMemoryCommandRegistry::new();
registry
.register_target("container-1", CommandTargetType::Container)
.await
.unwrap();
let result = resolved(®istry, None).await.unwrap();
assert_eq!(
result.target,
CommandTarget::new("container-1", CommandTargetType::Container)
);
}
#[tokio::test]
async fn test_resolve_shorthand_two_targets_is_ambiguous() {
let registry = InMemoryCommandRegistry::new();
registry
.register_target("worker-a", CommandTargetType::Worker)
.await
.unwrap();
registry
.register_target("worker-b", CommandTargetType::Worker)
.await
.unwrap();
let err = resolved(®istry, None).await.unwrap_err();
assert_eq!(err.code, "COMMAND_TARGET_AMBIGUOUS");
assert_eq!(err.http_status_code, Some(409));
}
#[tokio::test]
async fn test_resolve_shorthand_zero_targets_is_no_targets() {
let registry = InMemoryCommandRegistry::new();
let err = resolved(®istry, None).await.unwrap_err();
assert_eq!(err.code, "NO_COMMAND_TARGETS");
assert_eq!(err.http_status_code, Some(422));
}
#[tokio::test]
async fn test_delivery_mode_container_and_daemon_always_pull() {
let registry =
InMemoryCommandRegistry::with_worker_delivery_mode(CommandDeliveryMode::Push);
registry
.register_target("container-1", CommandTargetType::Container)
.await
.unwrap();
registry
.register_target("daemon-1", CommandTargetType::Daemon)
.await
.unwrap();
let container = resolved(®istry, Some("container-1")).await.unwrap();
assert_eq!(container.delivery_mode, CommandDeliveryMode::Pull);
let daemon = resolved(®istry, Some("daemon-1")).await.unwrap();
assert_eq!(daemon.delivery_mode, CommandDeliveryMode::Pull);
}
#[tokio::test]
async fn test_delivery_mode_worker_follows_registered_context() {
let push_registry =
InMemoryCommandRegistry::with_worker_delivery_mode(CommandDeliveryMode::Push);
push_registry
.register_target("worker-1", CommandTargetType::Worker)
.await
.unwrap();
let push_worker = push_registry
.resolve_target("dep-1", Some("worker-1"))
.await
.unwrap();
assert_eq!(push_worker.delivery_mode, CommandDeliveryMode::Push);
let pull_registry =
InMemoryCommandRegistry::with_worker_delivery_mode(CommandDeliveryMode::Pull);
pull_registry
.register_target("worker-1", CommandTargetType::Worker)
.await
.unwrap();
let pull_worker = pull_registry
.resolve_target("dep-1", Some("worker-1"))
.await
.unwrap();
assert_eq!(pull_worker.delivery_mode, CommandDeliveryMode::Pull);
}
#[tokio::test]
async fn test_create_command_stores_target_in_status_and_envelope_data() {
let registry = InMemoryCommandRegistry::new();
registry
.register_target("daemon-1", CommandTargetType::Daemon)
.await
.unwrap();
let resolved_target = registry.resolve_target("dep-1", None).await.unwrap();
let metadata = registry
.create_command(
"dep-1",
"sync-data",
&resolved_target,
CommandState::Pending,
None,
None,
)
.await
.unwrap();
let expected = CommandTarget::new("daemon-1", CommandTargetType::Daemon);
assert_eq!(metadata.target, expected);
assert_eq!(metadata.delivery_mode, CommandDeliveryMode::Pull);
let status = registry
.get_command_status(&metadata.command_id)
.await
.unwrap()
.unwrap();
assert_eq!(status.target, expected);
let envelope_data = registry
.get_command_metadata(&metadata.command_id)
.await
.unwrap()
.unwrap();
assert_eq!(envelope_data.target, expected);
assert_eq!(envelope_data.delivery_mode, CommandDeliveryMode::Pull);
}
#[tokio::test]
async fn non_terminal_update_cannot_resurrect_completed_command() {
let registry = InMemoryCommandRegistry::new();
registry
.register_target("daemon-1", CommandTargetType::Daemon)
.await
.unwrap();
let target = registry.resolve_target("dep-1", None).await.unwrap();
let command = registry
.create_command(
"dep-1",
"run",
&target,
CommandState::Dispatched,
None,
None,
)
.await
.unwrap();
assert!(registry
.complete_command(
&command.command_id,
CommandState::Succeeded,
Utc::now(),
None,
None,
)
.await
.unwrap());
assert!(!registry
.update_command_state(
&command.command_id,
CommandState::Pending,
None,
None,
None,
None,
)
.await
.unwrap());
assert_eq!(
registry
.get_command_status(&command.command_id)
.await
.unwrap()
.unwrap()
.state,
CommandState::Succeeded
);
}
#[tokio::test]
async fn test_register_target_rejects_colon_in_id() {
let registry = InMemoryCommandRegistry::new();
let err = registry
.register_target("evil:pending:x", CommandTargetType::Worker)
.await
.unwrap_err();
assert_eq!(err.code, "COMMAND_TARGET_ID_INVALID");
assert_eq!(err.http_status_code, Some(400));
}
#[tokio::test]
async fn test_resolve_explicit_colon_id_is_invalid() {
let registry = InMemoryCommandRegistry::new();
registry
.register_target("worker-a", CommandTargetType::Worker)
.await
.unwrap();
let err = resolved(®istry, Some("worker-a:pending:1"))
.await
.unwrap_err();
assert_eq!(err.code, "COMMAND_TARGET_ID_INVALID");
}
#[test]
fn test_select_command_target_rejects_registered_colon_id() {
let targets = vec![CommandTarget::new("a:pending:x", CommandTargetType::Worker)];
let err = select_command_target("dep-1", &targets, None).unwrap_err();
assert_eq!(err.code, "COMMAND_TARGET_ID_INVALID");
}
}