1use std::collections::BTreeSet;
4
5use serde::{Deserialize, Serialize};
6use serde_json::Value;
7
8use super::{
9 CapabilityBinding, CapabilityCardinality, CapabilityRequirementPlan, EventAdmissionPlan,
10 ExecutionClassId, ExecutionLanePlan, PLAN_SCHEMA_VERSION, PlanResolutionError,
11 PluginInstancePlan, RequestAdmissionPlan, ResolvedAppPlan, default_execution_lanes,
12};
13
14pub(crate) const fn old_authoring_version() -> u32 {
15 1
16}
17
18#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
20#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
21pub enum TerminalPolicy {
22 RequiredPath,
24 HostEssential {
27 roots: Vec<String>,
28 closure: Vec<String>,
29 },
30}
31
32impl TerminalPolicy {
33 pub(super) fn validate(
34 &self,
35 instances: &[PluginInstancePlan],
36 bindings: &[CapabilityBinding],
37 ) -> Result<(), PlanResolutionError> {
38 match self {
39 Self::RequiredPath => Ok(()),
40 Self::HostEssential { roots, closure } => {
41 validate_sorted_unique("roots", roots)?;
42 validate_sorted_unique("closure", closure)?;
43 let selected = instances
44 .iter()
45 .map(PluginInstancePlan::instance_key)
46 .collect::<BTreeSet<_>>();
47 for instance in roots.iter().chain(closure) {
48 if !selected.contains(instance.as_str()) {
49 return Err(PlanResolutionError::InvalidTerminalPolicy {
50 detail: format!("unknown Plugin Instance `{instance}`"),
51 });
52 }
53 }
54 let mut expected = roots.iter().cloned().collect::<BTreeSet<_>>();
55 let mut pending = roots.clone();
56 while let Some(consumer) = pending.pop() {
57 let Some(instance) = instances
58 .iter()
59 .find(|instance| instance.instance_key() == consumer)
60 else {
61 continue;
62 };
63 for requirement in
64 instance
65 .required_capabilities()
66 .iter()
67 .filter(|requirement| {
68 requirement.cardinality() == CapabilityCardinality::One
69 })
70 {
71 for provider in bindings.iter().filter(|binding| {
72 binding.consumer_instance() == consumer
73 && binding.requirement_id() == requirement.requirement_id()
74 }) {
75 if expected.insert(provider.provider_instance().to_owned()) {
76 pending.push(provider.provider_instance().to_owned());
77 }
78 }
79 }
80 }
81 let expected = expected.into_iter().collect::<Vec<_>>();
82 if *closure != expected {
83 return Err(PlanResolutionError::InvalidTerminalPolicy {
84 detail: format!(
85 "materialized closure {closure:?} does not match recomputed closure {expected:?}"
86 ),
87 });
88 }
89 Ok(())
90 }
91 }
92 }
93}
94
95fn validate_sorted_unique(field: &str, values: &[String]) -> Result<(), PlanResolutionError> {
96 if values.windows(2).any(|pair| pair[0] >= pair[1]) {
97 return Err(PlanResolutionError::InvalidTerminalPolicy {
98 detail: format!("{field} must be sorted and unique"),
99 });
100 }
101 Ok(())
102}
103
104#[derive(Deserialize)]
105#[serde(deny_unknown_fields)]
106pub(super) struct RequirementWire {
107 requirement_id: Option<String>,
108 capability_id: String,
109 descriptor_version: String,
110 cardinality: CapabilityCardinality,
111}
112
113impl From<RequirementWire> for CapabilityRequirementPlan {
114 fn from(wire: RequirementWire) -> Self {
115 Self {
116 requirement_id: wire
117 .requirement_id
118 .unwrap_or_else(|| format!("~{}", wire.capability_id)),
119 capability_id: wire.capability_id,
120 descriptor_version: wire.descriptor_version,
121 cardinality: wire.cardinality,
122 }
123 }
124}
125
126#[derive(Deserialize)]
127#[serde(deny_unknown_fields)]
128pub(super) struct BindingWire {
129 requirement_id: Option<String>,
130 consumer_instance: String,
131 capability_id: String,
132 descriptor_version: String,
133 provider_instance: String,
134 provider_order: usize,
135 admission: RequestAdmissionPlan,
136 admission_explicit: bool,
137 event_admission: EventAdmissionPlan,
138 event_admission_explicit: bool,
139}
140
141impl From<BindingWire> for CapabilityBinding {
142 fn from(wire: BindingWire) -> Self {
143 Self {
144 requirement_id: wire
145 .requirement_id
146 .unwrap_or_else(|| format!("~{}", wire.capability_id)),
147 consumer_instance: wire.consumer_instance,
148 capability_id: wire.capability_id,
149 descriptor_version: wire.descriptor_version,
150 provider_instance: wire.provider_instance,
151 provider_order: wire.provider_order,
152 admission: wire.admission,
153 admission_explicit: wire.admission_explicit,
154 event_admission: wire.event_admission,
155 event_admission_explicit: wire.event_admission_explicit,
156 }
157 }
158}
159
160#[derive(Deserialize)]
161#[serde(transparent)]
162pub(super) struct PlanWire(Value);
163
164#[derive(Deserialize)]
165#[serde(deny_unknown_fields)]
166struct DecodedPlan {
167 schema_version: u32,
168 terminal_policy: Option<TerminalPolicy>,
169 plugin_instances: Vec<PluginInstancePlan>,
170 capability_bindings: Vec<CapabilityBinding>,
171 #[serde(default = "default_execution_lanes")]
172 execution_lanes: Vec<ExecutionLanePlan>,
173}
174
175impl TryFrom<PlanWire> for ResolvedAppPlan {
176 type Error = String;
177
178 fn try_from(wire: PlanWire) -> Result<Self, String> {
179 let version = wire
180 .0
181 .get("schema_version")
182 .and_then(Value::as_u64)
183 .ok_or("missing Plan schema_version")?;
184 if version != 2 && version != u64::from(PLAN_SCHEMA_VERSION) {
185 return Err(format!("unsupported Plan schema version {version}"));
186 }
187 let modern = version == u64::from(PLAN_SCHEMA_VERSION);
188 check_field(&wire.0, "terminal_policy", modern)?;
189 for instance in wire
190 .0
191 .get("plugin_instances")
192 .and_then(Value::as_array)
193 .into_iter()
194 .flatten()
195 {
196 check_field(instance, "authoring_version", modern)?;
197 check_field(instance, "runtime_profile", modern)?;
198 for requirement in instance
199 .get("required_capabilities")
200 .and_then(Value::as_array)
201 .into_iter()
202 .flatten()
203 {
204 check_field(requirement, "requirement_id", modern)?;
205 }
206 }
207 for binding in wire
208 .0
209 .get("capability_bindings")
210 .and_then(Value::as_array)
211 .into_iter()
212 .flatten()
213 {
214 check_field(binding, "requirement_id", modern)?;
215 }
216 let mut decoded: DecodedPlan =
217 serde_json::from_value(wire.0).map_err(|error| error.to_string())?;
218 if decoded.schema_version == 2 {
219 for instance in &mut decoded.plugin_instances {
220 instance.runtime_profile = old_runtime_profile(&instance.execution_class);
221 }
222 }
223 Ok(Self {
224 schema_version: PLAN_SCHEMA_VERSION,
225 terminal_policy: decoded
226 .terminal_policy
227 .unwrap_or(TerminalPolicy::RequiredPath),
228 plugin_instances: decoded.plugin_instances,
229 capability_bindings: decoded.capability_bindings,
230 execution_lanes: decoded.execution_lanes,
231 })
232 }
233}
234
235fn check_field(value: &Value, field: &str, expected: bool) -> Result<(), String> {
236 if expected && value.get(field).is_some_and(Value::is_null) {
237 return Err(format!("Plan schema requires non-null {field}"));
238 }
239 if value.get(field).is_some() != expected {
240 return Err(format!(
241 "Plan schema {} {field}",
242 if expected { "requires" } else { "forbids" }
243 ));
244 }
245 Ok(())
246}
247
248pub(crate) fn old_runtime_profile(class: &ExecutionClassId) -> String {
249 match class.as_str() {
250 "lenso.native-rust@1" => "lenso.native-authoring@1".to_owned(),
251 other => other.to_owned(),
254 }
255}
256
257pub(super) fn validate_authoring(instance: &PluginInstancePlan) -> Result<(), PlanResolutionError> {
258 if !matches!(instance.authoring_version, 1 | 2) || instance.runtime_profile.trim().is_empty() {
259 return Err(PlanResolutionError::InvalidAuthoring {
260 instance_key: instance.instance_key.clone(),
261 });
262 }
263 for requirement in &instance.required_capabilities {
264 let id = requirement.requirement_id();
265 let valid = if instance.authoring_version == 1 {
266 id == format!("~{}", requirement.capability_id())
267 } else {
268 valid_requirement_id(id)
269 };
270 if !valid {
271 return Err(PlanResolutionError::InvalidRequirementId {
272 consumer_instance: instance.instance_key.clone(),
273 requirement_id: id.to_owned(),
274 });
275 }
276 }
277 Ok(())
278}
279
280pub(crate) fn valid_requirement_id(id: &str) -> bool {
281 (1..=64).contains(&id.len())
282 && id.as_bytes()[0].is_ascii_lowercase()
283 && id
284 .bytes()
285 .all(|byte| byte.is_ascii_lowercase() || byte.is_ascii_digit() || byte == b'_')
286}