use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::{
manifest::{CapabilityDeclarations, ProviderRole},
scope::{ScopeEnded, ScopeRecord, ScopeRecordResult, ScopeStamp, ScopeStatus},
BindIdentity, Principal, RouteCloseReason, RouteTarget,
};
pub const MODULE_CONTROL_OP_HEALTH_CHECK: &str = "health.check";
pub const MODULE_TO_SUBC_OP_CATALOG_UPDATE: &str = "catalog.update";
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum HealthStatus {
Ok,
Degraded,
Failing,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct HealthReport {
pub status: HealthStatus,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub detail: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub metrics: Option<Value>,
}
impl HealthReport {
pub fn ok() -> Self {
Self {
status: HealthStatus::Ok,
detail: None,
metrics: None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(tag = "op")]
#[allow(clippy::large_enum_variant)]
pub enum ModuleControlRequest {
#[serde(rename = "route.bind")]
RouteBind {
route_channel: u16,
epoch: u32,
target: RouteTarget,
identity: BindIdentity,
#[serde(default, skip_serializing_if = "Option::is_none")]
principal: Option<Principal>,
#[serde(default, skip_serializing_if = "Option::is_none")]
consumer_capabilities: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
role_versions: Option<BTreeMap<String, String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
admission_facts: Option<Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
scope: Option<ScopeStamp>,
},
#[serde(rename = "health.check")]
HealthCheck {},
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(tag = "op")]
pub enum ModuleControlCommand {
#[serde(rename = "module.draining")]
Draining {
reason: RouteCloseReason,
deadline_ms: u64,
},
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(tag = "op")]
pub enum ModuleControlResponse {
#[serde(rename = "route.bind")]
RouteBindAck {},
#[serde(rename = "health.check")]
HealthCheck {
status: HealthStatus,
#[serde(default, skip_serializing_if = "Option::is_none")]
detail: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
metrics: Option<Value>,
},
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(tag = "op")]
pub enum ModuleControlRequestFromModule {
#[serde(rename = "catalog.update")]
CatalogUpdate {
provides: Vec<ProviderRole>,
#[serde(default, skip_serializing_if = "Option::is_none")]
capabilities: Option<CapabilityDeclarations>,
#[serde(default, skip_serializing_if = "Option::is_none")]
ready: Option<bool>,
},
#[serde(rename = "supervisor.live_roots")]
LiveRoots {},
#[serde(rename = "scope.sync")]
ScopeSync {
generation: u64,
scopes: Vec<ScopeRecord>,
},
#[serde(rename = "scope.describe")]
ScopeDescribe {
owner: Principal,
#[serde(rename = "ref")]
scope_ref: String,
},
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct LiveRoot {
pub project_root: std::path::PathBuf,
pub bound: u64,
pub pending: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(tag = "op")]
pub enum ModuleControlResponseToModule {
#[serde(rename = "catalog.update")]
CatalogUpdate {},
#[serde(rename = "supervisor.live_roots")]
LiveRoots {
roots: Vec<LiveRoot>,
unknown_root_bindings: u64,
total_bindings: u64,
},
#[serde(rename = "scope.sync")]
ScopeSync {
generation: u64,
results: Vec<ScopeRecordResult>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
ended: Vec<ScopeEnded>,
},
#[serde(rename = "scope.describe")]
ScopeDescribe {
status: ScopeStatus,
#[serde(default, skip_serializing_if = "Option::is_none")]
scope_epoch: Option<u64>,
daemon_incarnation: String,
owner_synced: bool,
owner_configured: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
scope: Option<ScopeStamp>,
},
}
impl From<HealthReport> for ModuleControlResponse {
fn from(report: HealthReport) -> Self {
Self::HealthCheck {
status: report.status,
detail: report.detail,
metrics: report.metrics,
}
}
}
impl ModuleControlResponse {
pub fn health_report(&self) -> Option<HealthReport> {
match self {
Self::HealthCheck {
status,
detail,
metrics,
} => Some(HealthReport {
status: *status,
detail: detail.clone(),
metrics: metrics.clone(),
}),
Self::RouteBindAck {} => None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(tag = "op")]
pub enum ModuleControlPush {
#[serde(rename = "route.status")]
RouteStatus {
route_channel: u16,
route_epoch: u32,
status: String,
},
}
pub const ROLE_VERSIONS_FIELD: &str = "role_versions";
pub const MAX_ROLE_VERSIONS: usize = 8;
pub const MAX_ROLE_NAME_LEN: usize = 64;
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum RoleVersionsError {
TooMany { count: usize },
InvalidRole { role: String },
InvalidVersion { role: String, version: String },
}
impl RoleVersionsError {
pub fn field(&self) -> &'static str {
ROLE_VERSIONS_FIELD
}
}
impl std::fmt::Display for RoleVersionsError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::TooMany { count } => write!(
f,
"{ROLE_VERSIONS_FIELD} has {count} entries; at most {MAX_ROLE_VERSIONS} are allowed"
),
Self::InvalidRole { role } => write!(
f,
"{ROLE_VERSIONS_FIELD} names role {role:?}, which is not lowercase letters and \
digits in words joined by '-', at most {MAX_ROLE_NAME_LEN} bytes"
),
Self::InvalidVersion { role, version } => write!(
f,
"{ROLE_VERSIONS_FIELD} gives role {role:?} version {version:?}, which is not 'v' \
followed by a positive integer without leading zeros"
),
}
}
}
impl std::error::Error for RoleVersionsError {}
pub fn validate_role_versions(
role_versions: &BTreeMap<String, String>,
) -> Result<(), RoleVersionsError> {
if role_versions.len() > MAX_ROLE_VERSIONS {
return Err(RoleVersionsError::TooMany {
count: role_versions.len(),
});
}
for (role, version) in role_versions {
if !is_role_name(role) {
return Err(RoleVersionsError::InvalidRole { role: role.clone() });
}
if !is_role_version(version) {
return Err(RoleVersionsError::InvalidVersion {
role: role.clone(),
version: version.clone(),
});
}
}
Ok(())
}
fn is_role_name(role: &str) -> bool {
!role.is_empty()
&& role.len() <= MAX_ROLE_NAME_LEN
&& role.split('-').all(|word| {
!word.is_empty()
&& word
.bytes()
.all(|byte| byte.is_ascii_lowercase() || byte.is_ascii_digit())
})
}
fn is_role_version(version: &str) -> bool {
let bytes = version.as_bytes();
bytes.len() >= 2
&& bytes[0] == b'v'
&& (b'1'..=b'9').contains(&bytes[1])
&& bytes[2..].iter().all(u8::is_ascii_digit)
}
#[cfg(test)]
mod tests {
use super::*;
fn map(entries: &[(&str, &str)]) -> BTreeMap<String, String> {
entries
.iter()
.map(|(role, version)| (role.to_string(), version.to_string()))
.collect()
}
#[test]
fn role_versions_accept_role_names_and_positive_versions() {
for entries in [
vec![],
vec![("tool-provider", "v1")],
vec![("a", "v9"), ("b2", "v10"), ("x-1-y", "v1203")],
] {
assert_eq!(
validate_role_versions(&map(&entries)),
Ok(()),
"{entries:?}"
);
}
let longest = "a".repeat(MAX_ROLE_NAME_LEN);
assert_eq!(validate_role_versions(&map(&[(&longest, "v1")])), Ok(()));
let full: BTreeMap<String, String> = (0..MAX_ROLE_VERSIONS)
.map(|index| (format!("role-{index}"), "v1".to_string()))
.collect();
assert_eq!(validate_role_versions(&full), Ok(()));
}
#[test]
fn role_versions_refuse_malformed_role_names() {
let too_long = "a".repeat(MAX_ROLE_NAME_LEN + 1);
for role in [
"",
"Tool-provider",
"tool_provider",
"tool provider",
"-tool",
"tool-",
"tool--provider",
"tool.provider",
"outil-é",
too_long.as_str(),
] {
let error = validate_role_versions(&map(&[(role, "v1")])).unwrap_err();
assert_eq!(
error,
RoleVersionsError::InvalidRole {
role: role.to_string()
},
"{role:?}"
);
assert_eq!(error.field(), "role_versions");
}
}
#[test]
fn role_versions_refuse_malformed_versions() {
for version in [
"", "v", "v0", "v01", "1", "V1", "v1.0", "v-1", "v1 ", " v1", "vx",
] {
let error = validate_role_versions(&map(&[("tool-provider", version)])).unwrap_err();
assert_eq!(
error,
RoleVersionsError::InvalidVersion {
role: "tool-provider".to_string(),
version: version.to_string(),
},
"{version:?}"
);
assert_eq!(error.field(), "role_versions");
}
}
#[test]
fn role_versions_refuse_more_than_eight_entries() {
let nine: BTreeMap<String, String> = (0..=MAX_ROLE_VERSIONS)
.map(|index| (format!("role-{index}"), "v1".to_string()))
.collect();
let error = validate_role_versions(&nine).unwrap_err();
assert_eq!(error, RoleVersionsError::TooMany { count: 9 });
assert_eq!(error.field(), "role_versions");
assert!(error.to_string().starts_with("role_versions"), "{error}");
}
#[test]
fn route_bind_omits_absent_role_versions_and_carries_present_ones_verbatim() {
let bind = |role_versions| ModuleControlRequest::RouteBind {
route_channel: 1,
epoch: 1,
target: crate::RouteTarget::ToolProvider {
module_id: "aft".to_string(),
},
identity: crate::BindIdentity::new("/tmp/p", "h", "s"),
principal: None,
consumer_capabilities: None,
role_versions,
admission_facts: None,
scope: None,
};
let absent = serde_json::to_value(bind(None)).unwrap();
assert!(absent.get("role_versions").is_none(), "{absent}");
let present = bind(Some(map(&[("tool-provider", "v1")])));
let encoded = serde_json::to_value(&present).unwrap();
assert_eq!(
encoded["role_versions"],
serde_json::json!({ "tool-provider": "v1" })
);
let decoded: ModuleControlRequest = serde_json::from_value(encoded).unwrap();
assert_eq!(decoded, present);
}
}