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::binding::StreamFraming;
422    use crate::ir::contract::{
423        FieldIr, Idempotency, InputShape, MethodIr, MethodKind, PrimitiveType, ServiceIr, TypeRef,
424    };
425
426    fn one_method_contract() -> ContractIr {
427        ServiceIr {
428            name: "Svc".into(),
429            gear: "m".into(),
430            version: "v1".into(),
431            methods: vec![MethodIr {
432                name: "do_thing".into(),
433                kind: MethodKind::Unary,
434                input: InputShape {
435                    fields: vec![FieldIr {
436                        name: "id".into(),
437                        ty: TypeRef::Primitive(PrimitiveType::String),
438                        optional: false,
439                        role: crate::ir::contract::FieldRole::Wire,
440                    }],
441                },
442                output: TypeRef::Named("Out".into()),
443                error: None,
444                idempotency: Idempotency::SafeRead,
445                optional: false,
446            }],
447        }
448    }
449
450    #[test]
451    fn rejects_query_binding_to_unknown_field() {
452        let contract = one_method_contract();
453        let binding = HttpBindingIr {
454            base_path: "/api".into(),
455            methods: vec![HttpMethodBindingIr {
456                method_name: "do_thing".into(),
457                http_method: HttpMethod::Get,
458                path_template: "/things".into(),
459                field_bindings: vec![HttpFieldBinding::Query {
460                    field: "missing".into(),
461                    param: "missing".into(),
462                }],
463                retryable: false,
464                streaming: false,
465                stream_framing: StreamFraming::default(),
466                optional: false,
467            }],
468        };
469        let errs = validate_http_binding(&contract, &binding).unwrap_err();
470        assert!(
471            errs.iter().any(|e| e
472                .message
473                .contains("Query binding references field 'missing'")),
474            "expected query field-ref error, got: {errs:?}"
475        );
476    }
477
478    #[test]
479    fn rejects_duplicate_body_bindings() {
480        let contract = one_method_contract();
481        let binding = HttpBindingIr {
482            base_path: "/api".into(),
483            methods: vec![HttpMethodBindingIr {
484                method_name: "do_thing".into(),
485                http_method: HttpMethod::Post,
486                path_template: "/things".into(),
487                field_bindings: vec![HttpFieldBinding::Body, HttpFieldBinding::Body],
488                retryable: false,
489                streaming: false,
490                stream_framing: StreamFraming::default(),
491                optional: false,
492            }],
493        };
494        let errs = validate_http_binding(&contract, &binding).unwrap_err();
495        assert!(
496            errs.iter().any(|e| e.message.contains("Body bindings")),
497            "expected duplicate Body error, got: {errs:?}"
498        );
499    }
500
501    #[test]
502    fn rejects_base_path_without_leading_slash() {
503        let contract = one_method_contract();
504        let binding = HttpBindingIr {
505            base_path: "api/m/v1".into(),
506            methods: vec![HttpMethodBindingIr {
507                method_name: "do_thing".into(),
508                http_method: HttpMethod::Get,
509                path_template: "/things".into(),
510                field_bindings: vec![],
511                retryable: false,
512                streaming: false,
513                stream_framing: StreamFraming::default(),
514                optional: false,
515            }],
516        };
517        let errs = validate_http_binding(&contract, &binding).unwrap_err();
518        assert!(
519            errs.iter()
520                .any(|e| e.message.contains("base_path must start with '/'")),
521            "expected base_path slash error, got: {errs:?}"
522        );
523    }
524
525    #[test]
526    fn rejects_unbalanced_path_template_braces() {
527        let contract = one_method_contract();
528        let binding = HttpBindingIr {
529            base_path: "/api".into(),
530            methods: vec![HttpMethodBindingIr {
531                method_name: "do_thing".into(),
532                http_method: HttpMethod::Get,
533                path_template: "/things/{id".into(),
534                field_bindings: vec![HttpFieldBinding::Path {
535                    field: "id".into(),
536                    param: "id".into(),
537                }],
538                retryable: false,
539                streaming: false,
540                stream_framing: StreamFraming::default(),
541                optional: false,
542            }],
543        };
544        let errs = validate_http_binding(&contract, &binding).unwrap_err();
545        assert!(
546            errs.iter()
547                .any(|e| e.message.contains("unclosed '{'")
548                    || e.message.contains("unbalanced braces")),
549            "expected unbalanced brace error, got: {errs:?}"
550        );
551    }
552
553    #[test]
554    fn accepts_valid_binding() {
555        let contract = one_method_contract();
556        let binding = HttpBindingIr {
557            base_path: "/api".into(),
558            methods: vec![HttpMethodBindingIr {
559                method_name: "do_thing".into(),
560                http_method: HttpMethod::Get,
561                path_template: "/things/{id}".into(),
562                field_bindings: vec![HttpFieldBinding::Path {
563                    field: "id".into(),
564                    param: "id".into(),
565                }],
566                retryable: false,
567                streaming: false,
568                stream_framing: StreamFraming::default(),
569                optional: false,
570            }],
571        };
572        validate_http_binding(&contract, &binding).expect("valid binding should pass");
573    }
574}
575
576#[cfg(test)]
577#[cfg_attr(coverage_nightly, coverage(off))]
578mod version_base_path_tests {
579    use super::version_matches_base_path;
580
581    #[test]
582    fn accepts_matching_version_segment() {
583        assert!(version_matches_base_path("v1", "/api/billing/v1"));
584        assert!(version_matches_base_path("v2", "/api/api-contracts/v2"));
585        assert!(version_matches_base_path("v10", "/api/x/v10"));
586        // The version need not be the last path segment, only the last
587        // version-shaped one.
588        assert!(version_matches_base_path("v1", "/api/v1/payments"));
589    }
590
591    #[test]
592    fn rejects_disagreeing_version() {
593        assert!(!version_matches_base_path("v2", "/api/billing/v1"));
594        assert!(!version_matches_base_path("v2", "/api/billing"));
595        // A half-applied major bump: the version-designating (last) segment is
596        // still v1, so this must not pass just because "v2" appears somewhere.
597        assert!(!version_matches_base_path("v2", "/v2/api/billing/v1"));
598    }
599
600    #[test]
601    fn rejects_vacuous_and_malformed_inputs() {
602        // An empty version must never match — a leading-slash base_path splits
603        // into a leading empty segment, which would otherwise pass.
604        assert!(!version_matches_base_path("", "/api/billing/v1"));
605        assert!(!version_matches_base_path("", ""));
606        // `v` alone and non-numeric tails are not version segments.
607        assert!(!version_matches_base_path("v", "/api/v"));
608        assert!(!version_matches_base_path("v2", "/api/xv2"));
609        assert!(!version_matches_base_path("v2", "/api/v2beta"));
610    }
611}