Skip to main content

toolkit_contract/ir/
validation.rs

1use super::binding::{HttpBindingIr, HttpFieldBinding, HttpMethod, HttpMethodBindingIr};
2use super::contract::ContractIr;
3use std::collections::HashSet;
4use std::fmt;
5
6#[derive(Debug, Clone)]
7pub struct ValidationError {
8    pub location: String,
9    pub message: String,
10}
11
12impl fmt::Display for ValidationError {
13    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
14        write!(f, "{}: {}", self.location, self.message)
15    }
16}
17
18/// Whether a projection's `base_path` spells the same major version the
19/// contract declares (ADR-0007 §3).
20///
21/// The rule is the **last version-shaped segment wins**: a `base_path` may carry
22/// other segments, but the version-designating one — the last segment matching
23/// `v<digits>` — must equal `version`. That rejects a half-applied major bump
24/// such as `version = "v2"` against `base_path = "/v2/api/billing/v1"`, which a
25/// "contains the version anywhere" check would wave through.
26///
27/// An empty `version`, or a `base_path` with no version-shaped segment at all,
28/// is never a match — otherwise the check would pass vacuously (a leading-slash
29/// path splits into a leading empty segment).
30///
31/// Used by the `require_full_coverage` assertion the REST projection macro
32/// generates; exposed here so it is unit-testable and so generated code has a
33/// single definition to call.
34#[must_use]
35pub fn version_matches_base_path(version: &str, base_path: &str) -> bool {
36    if version.is_empty() {
37        return false;
38    }
39    let is_version_segment = |seg: &&str| {
40        seg.strip_prefix('v')
41            .is_some_and(|rest| !rest.is_empty() && rest.bytes().all(|b| b.is_ascii_digit()))
42    };
43    base_path
44        .split('/')
45        .rfind(is_version_segment)
46        .is_some_and(|seg| seg == version)
47}
48
49/// Validates a [`ContractIr`] for structural well-formedness.
50///
51/// # Errors
52/// Returns a vector of [`ValidationError`] describing every problem found
53/// (empty name/module/version, no methods, empty or duplicate method names).
54pub fn validate_contract(ir: &ContractIr) -> Result<(), Vec<ValidationError>> {
55    let mut errors = Vec::new();
56
57    if ir.name.is_empty() {
58        errors.push(ValidationError {
59            location: "ContractIr".to_owned(),
60            message: "contract name must not be empty".to_owned(),
61        });
62    }
63
64    if ir.gear.is_empty() {
65        errors.push(ValidationError {
66            location: "ContractIr".to_owned(),
67            message: "gear must not be empty".to_owned(),
68        });
69    }
70
71    if ir.version.is_empty() {
72        errors.push(ValidationError {
73            location: "ContractIr".to_owned(),
74            message: "version must not be empty".to_owned(),
75        });
76    }
77
78    if ir.methods.is_empty() {
79        errors.push(ValidationError {
80            location: "ContractIr".to_owned(),
81            message: "must have at least one method".to_owned(),
82        });
83    }
84
85    let mut seen_names: HashSet<&str> = HashSet::new();
86    for method in &ir.methods {
87        if method.name.is_empty() {
88            errors.push(ValidationError {
89                location: format!("ContractIr.methods[{}]", method.name),
90                message: "method name must not be empty".to_owned(),
91            });
92        } else if !seen_names.insert(&method.name) {
93            errors.push(ValidationError {
94                location: format!("ContractIr.methods[{0}]", method.name),
95                message: format!("duplicate method name: {}", method.name),
96            });
97        }
98    }
99
100    if errors.is_empty() {
101        Ok(())
102    } else {
103        Err(errors)
104    }
105}
106
107/// Validates an [`HttpBindingIr`] against its [`ContractIr`].
108///
109/// # Errors
110/// Returns a vector of [`ValidationError`] when `base_path` is malformed or method
111/// coverage / path-template / field-binding checks fail.
112pub fn validate_http_binding(
113    contract: &ContractIr,
114    binding: &HttpBindingIr,
115) -> Result<(), Vec<ValidationError>> {
116    let mut errors = Vec::new();
117
118    if binding.base_path.is_empty() {
119        errors.push(ValidationError {
120            location: "HttpBindingIr".to_owned(),
121            message: "base_path must not be empty".to_owned(),
122        });
123    } else if !binding.base_path.starts_with('/') {
124        errors.push(ValidationError {
125            location: "HttpBindingIr".to_owned(),
126            message: format!("base_path must start with '/': got '{}'", binding.base_path),
127        });
128    }
129
130    validate_method_coverage(contract, binding, &mut errors);
131
132    for method_binding in &binding.methods {
133        validate_single_method_binding(contract, method_binding, &mut errors);
134    }
135
136    if errors.is_empty() {
137        Ok(())
138    } else {
139        Err(errors)
140    }
141}
142
143fn validate_method_coverage(
144    contract: &ContractIr,
145    binding: &HttpBindingIr,
146    errors: &mut Vec<ValidationError>,
147) {
148    let contract_method_names: HashSet<&str> =
149        contract.methods.iter().map(|m| m.name.as_str()).collect();
150    let mut binding_method_names: HashSet<&str> = HashSet::new();
151
152    for method in &binding.methods {
153        let name = method.method_name.as_str();
154        if !binding_method_names.insert(name) {
155            errors.push(ValidationError {
156                location: format!("HttpBindingIr.methods[{name}]"),
157                message: format!("duplicate binding for contract method: {name}"),
158            });
159        }
160    }
161
162    for name in &contract_method_names {
163        if !binding_method_names.contains(name) {
164            errors.push(ValidationError {
165                location: format!("HttpBindingIr.methods[{name}]"),
166                message: format!("missing binding for contract method: {name}"),
167            });
168        }
169    }
170
171    for name in &binding_method_names {
172        if !contract_method_names.contains(name) {
173            errors.push(ValidationError {
174                location: format!("HttpBindingIr.methods[{name}]"),
175                message: format!("binding for unknown method not in contract: {name}"),
176            });
177        }
178    }
179}
180
181fn validate_single_method_binding(
182    contract: &ContractIr,
183    method_binding: &HttpMethodBindingIr,
184    errors: &mut Vec<ValidationError>,
185) {
186    let method_loc = format!("HttpBindingIr.methods[{}]", method_binding.method_name);
187
188    validate_body_constraint(method_binding, &method_loc, errors);
189    validate_path_params(method_binding, &method_loc, errors);
190    validate_path_template_braces(method_binding, &method_loc, errors);
191    validate_field_references(contract, method_binding, &method_loc, errors);
192    validate_single_body_binding(method_binding, &method_loc, errors);
193}
194
195fn validate_body_constraint(
196    method_binding: &HttpMethodBindingIr,
197    method_loc: &str,
198    errors: &mut Vec<ValidationError>,
199) {
200    if !matches!(
201        method_binding.http_method,
202        HttpMethod::Get | HttpMethod::Delete
203    ) {
204        return;
205    }
206
207    let has_body = method_binding
208        .field_bindings
209        .iter()
210        .any(|fb| matches!(fb, HttpFieldBinding::Body));
211
212    if has_body {
213        let verb = match method_binding.http_method {
214            HttpMethod::Get => "GET",
215            HttpMethod::Post => "POST",
216            HttpMethod::Put => "PUT",
217            HttpMethod::Patch => "PATCH",
218            HttpMethod::Delete => "DELETE",
219        };
220        errors.push(ValidationError {
221            location: method_loc.to_owned(),
222            message: format!("{verb} method must not have Body field binding"),
223        });
224    }
225}
226
227fn validate_path_params(
228    method_binding: &HttpMethodBindingIr,
229    method_loc: &str,
230    errors: &mut Vec<ValidationError>,
231) {
232    let template_params = extract_path_params(&method_binding.path_template);
233    let path_binding_params: HashSet<&str> = method_binding
234        .field_bindings
235        .iter()
236        .filter_map(|fb| {
237            if let HttpFieldBinding::Path { param, .. } = fb {
238                Some(param.as_str())
239            } else {
240                None
241            }
242        })
243        .collect();
244
245    for param in &template_params {
246        if !path_binding_params.contains(param.as_str()) {
247            errors.push(ValidationError {
248                location: method_loc.to_owned(),
249                message: format!(
250                    "path template parameter '{{{param}}}' has no corresponding Path field binding"
251                ),
252            });
253        }
254    }
255}
256
257fn validate_field_references(
258    contract: &ContractIr,
259    method_binding: &HttpMethodBindingIr,
260    method_loc: &str,
261    errors: &mut Vec<ValidationError>,
262) {
263    let Some(contract_method) = contract
264        .methods
265        .iter()
266        .find(|m| m.name == method_binding.method_name)
267    else {
268        return;
269    };
270
271    let input_field_names: HashSet<&str> = contract_method
272        .input
273        .fields
274        .iter()
275        .map(|f| f.name.as_str())
276        .collect();
277
278    for fb in &method_binding.field_bindings {
279        let (kind, field) = match fb {
280            HttpFieldBinding::Path { field, .. } => ("Path", field),
281            HttpFieldBinding::Query { field, .. } => ("Query", field),
282            HttpFieldBinding::Body => continue,
283        };
284        if !input_field_names.contains(field.as_str()) {
285            errors.push(ValidationError {
286                location: method_loc.to_owned(),
287                message: format!(
288                    "{kind} binding references field '{field}' not found in contract method input"
289                ),
290            });
291        }
292    }
293}
294
295fn validate_single_body_binding(
296    method_binding: &HttpMethodBindingIr,
297    method_loc: &str,
298    errors: &mut Vec<ValidationError>,
299) {
300    let body_count = method_binding
301        .field_bindings
302        .iter()
303        .filter(|fb| matches!(fb, HttpFieldBinding::Body))
304        .count();
305    if body_count > 1 {
306        errors.push(ValidationError {
307            location: method_loc.to_owned(),
308            message: format!(
309                "method has {body_count} Body bindings; at most one Body binding is allowed"
310            ),
311        });
312    }
313}
314
315fn validate_path_template_braces(
316    method_binding: &HttpMethodBindingIr,
317    method_loc: &str,
318    errors: &mut Vec<ValidationError>,
319) {
320    let template = &method_binding.path_template;
321    let mut depth = 0i32;
322    let mut current_param = String::new();
323    let mut in_param = false;
324    for ch in template.chars() {
325        match ch {
326            '{' => {
327                if in_param {
328                    errors.push(ValidationError {
329                        location: method_loc.to_owned(),
330                        message: format!(
331                            "path template '{template}' has nested '{{' before matching '}}'"
332                        ),
333                    });
334                    return;
335                }
336                in_param = true;
337                depth += 1;
338                current_param.clear();
339            }
340            '}' => {
341                if !in_param {
342                    errors.push(ValidationError {
343                        location: method_loc.to_owned(),
344                        message: format!("path template '{template}' has unmatched '}}'"),
345                    });
346                    return;
347                }
348                if current_param.is_empty() {
349                    errors.push(ValidationError {
350                        location: method_loc.to_owned(),
351                        message: format!("path template '{template}' has empty parameter '{{}}'"),
352                    });
353                }
354                if !current_param
355                    .chars()
356                    .all(|c| c.is_ascii_alphanumeric() || c == '_')
357                    || current_param
358                        .chars()
359                        .next()
360                        .is_some_and(|c| c.is_ascii_digit())
361                {
362                    errors.push(ValidationError {
363                        location: method_loc.to_owned(),
364                        message: format!(
365                            "path template '{template}' parameter '{{{current_param}}}' is not a valid identifier"
366                        ),
367                    });
368                }
369                in_param = false;
370                depth -= 1;
371                current_param.clear();
372            }
373            '/' => {
374                if in_param {
375                    errors.push(ValidationError {
376                        location: method_loc.to_owned(),
377                        message: format!(
378                            "path template '{template}' has unclosed '{{' before path separator"
379                        ),
380                    });
381                    return;
382                }
383            }
384            other => {
385                if in_param {
386                    current_param.push(other);
387                }
388            }
389        }
390    }
391    if depth != 0 {
392        errors.push(ValidationError {
393            location: method_loc.to_owned(),
394            message: format!("path template '{template}' has unbalanced braces"),
395        });
396    }
397}
398
399fn extract_path_params(template: &str) -> Vec<String> {
400    let mut params = Vec::new();
401    let mut rest = template;
402    while let Some(start) = rest.find('{') {
403        if let Some(end) = rest[start..].find('}') {
404            let param = &rest[start + 1..start + end];
405            if !param.is_empty() {
406                params.push(param.to_owned());
407            }
408            rest = &rest[start + end + 1..];
409        } else {
410            break;
411        }
412    }
413    params
414}
415
416#[cfg(test)]
417#[cfg_attr(coverage_nightly, coverage(off))]
418#[allow(clippy::unwrap_used)]
419mod tests {
420    use super::*;
421    use crate::ir::contract::{
422        FieldIr, Idempotency, InputShape, MethodIr, MethodKind, PrimitiveType, ServiceIr, TypeRef,
423    };
424
425    fn one_method_contract() -> ContractIr {
426        ServiceIr {
427            name: "Svc".into(),
428            gear: "m".into(),
429            version: "v1".into(),
430            methods: vec![MethodIr {
431                name: "do_thing".into(),
432                kind: MethodKind::Unary,
433                input: InputShape {
434                    fields: vec![FieldIr {
435                        name: "id".into(),
436                        ty: TypeRef::Primitive(PrimitiveType::String),
437                        optional: false,
438                        role: crate::ir::contract::FieldRole::Wire,
439                    }],
440                },
441                output: TypeRef::Named("Out".into()),
442                error: None,
443                idempotency: Idempotency::SafeRead,
444                optional: false,
445            }],
446        }
447    }
448
449    #[test]
450    fn rejects_query_binding_to_unknown_field() {
451        let contract = one_method_contract();
452        let binding = HttpBindingIr {
453            base_path: "/api".into(),
454            methods: vec![HttpMethodBindingIr {
455                method_name: "do_thing".into(),
456                http_method: HttpMethod::Get,
457                path_template: "/things".into(),
458                field_bindings: vec![HttpFieldBinding::Query {
459                    field: "missing".into(),
460                    param: "missing".into(),
461                }],
462                retryable: false,
463                streaming: false,
464                optional: false,
465            }],
466        };
467        let errs = validate_http_binding(&contract, &binding).unwrap_err();
468        assert!(
469            errs.iter().any(|e| e
470                .message
471                .contains("Query binding references field 'missing'")),
472            "expected query field-ref error, got: {errs:?}"
473        );
474    }
475
476    #[test]
477    fn rejects_duplicate_body_bindings() {
478        let contract = one_method_contract();
479        let binding = HttpBindingIr {
480            base_path: "/api".into(),
481            methods: vec![HttpMethodBindingIr {
482                method_name: "do_thing".into(),
483                http_method: HttpMethod::Post,
484                path_template: "/things".into(),
485                field_bindings: vec![HttpFieldBinding::Body, HttpFieldBinding::Body],
486                retryable: false,
487                streaming: false,
488                optional: false,
489            }],
490        };
491        let errs = validate_http_binding(&contract, &binding).unwrap_err();
492        assert!(
493            errs.iter().any(|e| e.message.contains("Body bindings")),
494            "expected duplicate Body error, got: {errs:?}"
495        );
496    }
497
498    #[test]
499    fn rejects_base_path_without_leading_slash() {
500        let contract = one_method_contract();
501        let binding = HttpBindingIr {
502            base_path: "api/m/v1".into(),
503            methods: vec![HttpMethodBindingIr {
504                method_name: "do_thing".into(),
505                http_method: HttpMethod::Get,
506                path_template: "/things".into(),
507                field_bindings: vec![],
508                retryable: false,
509                streaming: false,
510                optional: false,
511            }],
512        };
513        let errs = validate_http_binding(&contract, &binding).unwrap_err();
514        assert!(
515            errs.iter()
516                .any(|e| e.message.contains("base_path must start with '/'")),
517            "expected base_path slash error, got: {errs:?}"
518        );
519    }
520
521    #[test]
522    fn rejects_unbalanced_path_template_braces() {
523        let contract = one_method_contract();
524        let binding = HttpBindingIr {
525            base_path: "/api".into(),
526            methods: vec![HttpMethodBindingIr {
527                method_name: "do_thing".into(),
528                http_method: HttpMethod::Get,
529                path_template: "/things/{id".into(),
530                field_bindings: vec![HttpFieldBinding::Path {
531                    field: "id".into(),
532                    param: "id".into(),
533                }],
534                retryable: false,
535                streaming: false,
536                optional: false,
537            }],
538        };
539        let errs = validate_http_binding(&contract, &binding).unwrap_err();
540        assert!(
541            errs.iter()
542                .any(|e| e.message.contains("unclosed '{'")
543                    || e.message.contains("unbalanced braces")),
544            "expected unbalanced brace error, got: {errs:?}"
545        );
546    }
547
548    #[test]
549    fn accepts_valid_binding() {
550        let contract = one_method_contract();
551        let binding = HttpBindingIr {
552            base_path: "/api".into(),
553            methods: vec![HttpMethodBindingIr {
554                method_name: "do_thing".into(),
555                http_method: HttpMethod::Get,
556                path_template: "/things/{id}".into(),
557                field_bindings: vec![HttpFieldBinding::Path {
558                    field: "id".into(),
559                    param: "id".into(),
560                }],
561                retryable: false,
562                streaming: false,
563                optional: false,
564            }],
565        };
566        validate_http_binding(&contract, &binding).expect("valid binding should pass");
567    }
568}
569
570#[cfg(test)]
571#[cfg_attr(coverage_nightly, coverage(off))]
572mod version_base_path_tests {
573    use super::version_matches_base_path;
574
575    #[test]
576    fn accepts_matching_version_segment() {
577        assert!(version_matches_base_path("v1", "/api/billing/v1"));
578        assert!(version_matches_base_path("v2", "/api/api-contracts/v2"));
579        assert!(version_matches_base_path("v10", "/api/x/v10"));
580        // The version need not be the last path segment, only the last
581        // version-shaped one.
582        assert!(version_matches_base_path("v1", "/api/v1/payments"));
583    }
584
585    #[test]
586    fn rejects_disagreeing_version() {
587        assert!(!version_matches_base_path("v2", "/api/billing/v1"));
588        assert!(!version_matches_base_path("v2", "/api/billing"));
589        // A half-applied major bump: the version-designating (last) segment is
590        // still v1, so this must not pass just because "v2" appears somewhere.
591        assert!(!version_matches_base_path("v2", "/v2/api/billing/v1"));
592    }
593
594    #[test]
595    fn rejects_vacuous_and_malformed_inputs() {
596        // An empty version must never match — a leading-slash base_path splits
597        // into a leading empty segment, which would otherwise pass.
598        assert!(!version_matches_base_path("", "/api/billing/v1"));
599        assert!(!version_matches_base_path("", ""));
600        // `v` alone and non-numeric tails are not version segments.
601        assert!(!version_matches_base_path("v", "/api/v"));
602        assert!(!version_matches_base_path("v2", "/api/xv2"));
603        assert!(!version_matches_base_path("v2", "/api/v2beta"));
604    }
605}