Skip to main content

apif_utils/
coverage.rs

1//! gRPC method and protobuf message field coverage collector.
2//!
3//! Tracks which gRPC service/method calls were made during test execution
4//! and which protobuf message fields were covered by assertions.
5
6use prost_reflect::{DescriptorPool, MessageDescriptor};
7use serde::{Deserialize, Serialize};
8use std::collections::{HashMap, HashSet};
9use std::sync::{Arc, Mutex};
10
11/// Coverage data for a single file.
12#[derive(Debug, Clone, Serialize, Deserialize)]
13pub struct CoverageFile {
14    pub uri: String,
15    pub statements: CoverageStats,
16    #[serde(skip_serializing_if = "Option::is_none")]
17    pub branches: Option<CoverageStats>,
18    #[serde(skip_serializing_if = "Option::is_none")]
19    pub functions: Option<CoverageStats>,
20    #[serde(skip_serializing_if = "Option::is_none")]
21    pub fields: Option<CoverageStats>,
22}
23
24/// Coverage statistics (covered vs total).
25#[derive(Debug, Clone, Serialize, Deserialize)]
26pub struct CoverageStats {
27    pub covered: usize,
28    pub total: usize,
29}
30
31/// Coverage data for a protobuf message type's fields.
32#[derive(Debug, Clone, Serialize, Deserialize)]
33pub struct MessageFieldCoverage {
34    pub message_type: String,
35    pub covered_fields: Vec<String>,
36    pub total_fields: usize,
37}
38
39/// Full coverage report with file and message-level statistics.
40#[derive(Debug, Clone, Serialize, Deserialize)]
41pub struct CoverageReport {
42    pub files: Vec<CoverageFile>,
43    pub messages: Vec<MessageFieldCoverage>,
44    pub summary: CoverageStats,
45    pub field_summary: CoverageStats,
46}
47
48/// Collects gRPC method call and protobuf field coverage during test execution.
49#[derive(Debug, Clone)]
50pub struct CoverageCollector {
51    calls: Arc<Mutex<HashMap<String, HashMap<String, u64>>>>,
52    pool: Arc<Mutex<DescriptorPool>>,
53    fields_covered: Arc<Mutex<HashMap<String, HashSet<String>>>>,
54}
55
56impl CoverageCollector {
57    pub fn new() -> Self {
58        Self {
59            calls: Arc::new(Mutex::new(HashMap::new())),
60            pool: Arc::new(Mutex::new(DescriptorPool::new())),
61            fields_covered: Arc::new(Mutex::new(HashMap::new())),
62        }
63    }
64
65    pub fn record_call(&self, service: &str, method: &str) {
66        let mut calls = self.calls.lock().unwrap_or_else(|e| e.into_inner());
67        let service_calls = calls.entry(service.to_string()).or_default();
68        *service_calls.entry(method.to_string()).or_insert(0) += 1;
69    }
70
71    pub fn record_fields_from_json(&self, message_type: &str, json: &serde_json::Value) {
72        let mut fields = self
73            .fields_covered
74            .lock()
75            .unwrap_or_else(|e| e.into_inner());
76        let message_fields = fields.entry(message_type.to_string()).or_default();
77        Self::extract_fields_from_json(json, message_fields, "");
78    }
79
80    fn extract_fields_from_json(
81        json: &serde_json::Value,
82        fields: &mut HashSet<String>,
83        prefix: &str,
84    ) {
85        if let serde_json::Value::Object(map) = json {
86            for (key, value) in map {
87                let field_path = if prefix.is_empty() {
88                    key.clone()
89                } else {
90                    format!("{}.{}", prefix, key)
91                };
92                fields.insert(field_path.clone());
93                Self::extract_fields_from_json(value, fields, &field_path);
94            }
95        } else if let serde_json::Value::Array(arr) = json {
96            for item in arr {
97                Self::extract_fields_from_json(item, fields, prefix);
98            }
99        }
100    }
101
102    pub fn register_pool(&self, other: &DescriptorPool) {
103        let mut pool = self.pool.lock().unwrap_or_else(|e| e.into_inner());
104        for file in other.files() {
105            let _ = pool.add_file_descriptor_proto(file.file_descriptor_proto().clone());
106        }
107    }
108
109    fn count_message_fields(pool: &DescriptorPool, message_type: &str) -> usize {
110        if let Some(msg) = pool.get_message_by_name(message_type) {
111            Self::count_fields_recursive(&msg)
112        } else {
113            0
114        }
115    }
116
117    /// Count fields the way `extract_fields_from_json` records covered ones:
118    /// every field contributes its own dotted path, and message-typed fields
119    /// additionally contribute the nested paths of their sub-message. Without
120    /// this the denominator only counted top-level fields while the numerator
121    /// counted nested `parent.child` paths, understating the total.
122    fn count_fields_recursive(msg: &MessageDescriptor) -> usize {
123        fn count(msg: &MessageDescriptor, visited: &mut HashSet<String>) -> usize {
124            // Guard against recursive message types (e.g. a tree node whose
125            // field points back at its own type) to avoid unbounded recursion.
126            if !visited.insert(msg.full_name().to_string()) {
127                return 0;
128            }
129            let mut total = 0;
130            for field in msg.fields() {
131                total += 1;
132                // Recurse into nested messages, but not map entries: a map's
133                // keys are dynamic and can't be enumerated from the schema.
134                if !field.is_map()
135                    && let prost_reflect::Kind::Message(sub) = field.kind()
136                {
137                    total += count(&sub, visited);
138                }
139            }
140            visited.remove(msg.full_name());
141            total
142        }
143        count(msg, &mut HashSet::new())
144    }
145
146    pub fn generate_json_report(&self) -> CoverageReport {
147        let calls = self.calls.lock().unwrap_or_else(|e| e.into_inner());
148        let pool = self.pool.lock().unwrap_or_else(|e| e.into_inner());
149        let fields_covered = self
150            .fields_covered
151            .lock()
152            .unwrap_or_else(|e| e.into_inner());
153
154        let mut files = Vec::new();
155        let mut messages = Vec::new();
156        let mut total_covered = 0;
157        let mut total_methods = 0;
158        let mut total_fields_covered = 0;
159        let mut total_fields = 0;
160
161        // Method coverage - deduplicated iteration pattern
162        let mut services: Vec<_> = pool.services().collect();
163        services.sort_by(|a, b| a.name().cmp(b.name()));
164
165        for service in services {
166            let service_name = service.name();
167            if service_name.contains("reflection") {
168                continue;
169            }
170
171            let methods: Vec<_> = service.methods().collect();
172            // Calls are recorded under the fully-qualified service name
173            // ("package.Service", see runner.rs). Look up with the same
174            // FQN so services inside a proto `package` aren't reported 0%.
175            let called_methods = calls.get(service.full_name()).cloned().unwrap_or_default();
176
177            let covered = methods
178                .iter()
179                .filter(|m| called_methods.get(m.name()).unwrap_or(&0) > &0)
180                .count();
181            let total = methods.len();
182
183            if total > 0 {
184                total_covered += covered;
185                total_methods += total;
186
187                files.push(CoverageFile {
188                    uri: format!("grpc://{}", service_name),
189                    statements: CoverageStats { covered, total },
190                    branches: None,
191                    functions: Some(CoverageStats { covered, total }),
192                    fields: None,
193                });
194            }
195        }
196
197        // Message field coverage
198        let mut all_message_types: HashSet<String> = HashSet::new();
199        for message_type in fields_covered.keys() {
200            all_message_types.insert(message_type.clone());
201        }
202
203        let mut sorted_messages: Vec<_> = all_message_types.into_iter().collect();
204        sorted_messages.sort();
205
206        for message_type in sorted_messages {
207            let covered = fields_covered
208                .get(&message_type)
209                .map(|s| s.len())
210                .unwrap_or(0);
211            let total = Self::count_message_fields(&pool, &message_type);
212
213            if total > 0 {
214                total_fields_covered += covered.min(total);
215                total_fields += total;
216
217                let covered_fields: Vec<String> = fields_covered
218                    .get(&message_type)
219                    .map(|s| {
220                        let mut v: Vec<_> = s.iter().cloned().collect();
221                        v.sort();
222                        v
223                    })
224                    .unwrap_or_default();
225
226                messages.push(MessageFieldCoverage {
227                    message_type,
228                    covered_fields,
229                    total_fields: total,
230                });
231            }
232        }
233
234        CoverageReport {
235            files,
236            messages,
237            summary: CoverageStats {
238                covered: total_covered,
239                total: total_methods,
240            },
241            field_summary: CoverageStats {
242                covered: total_fields_covered,
243                total: total_fields,
244            },
245        }
246    }
247
248    pub fn generate_text_report(&self) -> String {
249        let calls = self.calls.lock().unwrap_or_else(|e| e.into_inner());
250        let pool = self.pool.lock().unwrap_or_else(|e| e.into_inner());
251        let fields_covered = self
252            .fields_covered
253            .lock()
254            .unwrap_or_else(|e| e.into_inner());
255
256        let mut report = String::new();
257        report.push_str("--- gRPC API Coverage Report ---\n\n");
258
259        // Method coverage
260        let mut services: Vec<_> = pool.services().collect();
261        services.sort_by(|a, b| a.name().cmp(b.name()));
262
263        if services.is_empty() {
264            report.push_str("No services found in descriptors.\n");
265            return report;
266        }
267
268        for service in services {
269            let service_name = service.name();
270            if service_name == "grpc.reflection.v1alpha.ServerReflection"
271                || service_name == "grpc.reflection.v1.ServerReflection"
272            {
273                continue;
274            }
275
276            report.push_str(&format!("Service: {}\n", service_name));
277
278            // Match the fully-qualified name used when recording calls.
279            let called_methods = calls.get(service.full_name()).cloned().unwrap_or_default();
280
281            let mut methods: Vec<_> = service.methods().collect();
282            methods.sort_by(|a, b| a.name().cmp(b.name()));
283
284            let mut covered_count = 0;
285            let total_count = methods.len();
286
287            for method in methods {
288                let method_name = method.name();
289                let count = called_methods.get(method_name).unwrap_or(&0);
290
291                let status = if *count > 0 {
292                    covered_count += 1;
293                    format!("✅ ({} calls)", count)
294                } else {
295                    "❌ (0 calls)".to_string()
296                };
297
298                report.push_str(&format!("  - {}: {}\n", method_name, status));
299            }
300
301            let coverage_pct = if total_count > 0 {
302                (covered_count as f64 / total_count as f64) * 100.0
303            } else {
304                0.0
305            };
306
307            report.push_str(&format!(
308                "  Coverage: {:.1}% ({}/{})\n\n",
309                coverage_pct, covered_count, total_count
310            ));
311        }
312
313        // Message field coverage
314        if !fields_covered.is_empty() {
315            report.push_str("--- Message Field Coverage ---\n\n");
316
317            let mut message_types: Vec<_> = fields_covered.keys().cloned().collect();
318            message_types.sort();
319
320            for message_type in message_types {
321                let covered = fields_covered
322                    .get(&message_type)
323                    .map(|s| s.len())
324                    .unwrap_or(0);
325                let total = Self::count_message_fields(&pool, &message_type);
326
327                if total > 0 {
328                    let pct = (covered.min(total) as f64 / total as f64) * 100.0;
329                    let status = if pct >= 100.0 {
330                        "✅"
331                    } else if pct > 0.0 {
332                        "⚠️"
333                    } else {
334                        "❌"
335                    };
336                    report.push_str(&format!(
337                        "{} {} ({}/{})\n",
338                        status,
339                        message_type,
340                        covered.min(total),
341                        total
342                    ));
343                }
344            }
345        }
346
347        report
348    }
349}
350
351impl Default for CoverageCollector {
352    fn default() -> Self {
353        Self::new()
354    }
355}
356
357#[cfg(test)]
358mod tests {
359    use super::*;
360    use prost_reflect::prost_types::{
361        DescriptorProto, FileDescriptorProto, MethodDescriptorProto, ServiceDescriptorProto,
362    };
363
364    /// Build a pool with a service inside a proto `package`, so its
365    /// fully-qualified name ("my.pkg.Greeter") differs from its short name.
366    fn pool_with_packaged_service() -> DescriptorPool {
367        let mut pool = DescriptorPool::new();
368        let file = FileDescriptorProto {
369            name: Some("test.proto".to_string()),
370            package: Some("my.pkg".to_string()),
371            message_type: vec![DescriptorProto {
372                name: Some("Empty".to_string()),
373                ..Default::default()
374            }],
375            service: vec![ServiceDescriptorProto {
376                name: Some("Greeter".to_string()),
377                method: vec![MethodDescriptorProto {
378                    name: Some("SayHello".to_string()),
379                    input_type: Some(".my.pkg.Empty".to_string()),
380                    output_type: Some(".my.pkg.Empty".to_string()),
381                    ..Default::default()
382                }],
383                ..Default::default()
384            }],
385            ..Default::default()
386        };
387        pool.add_file_descriptor_proto(file).unwrap();
388        pool
389    }
390
391    // Bug 6: calls are recorded under the fully-qualified service name, so
392    // coverage lookup must use the same FQN or packaged services report 0%.
393    #[test]
394    fn coverage_matches_fully_qualified_service_name() {
395        let collector = CoverageCollector::new();
396        collector.register_pool(&pool_with_packaged_service());
397        // Recorded exactly as runner.rs does: "package.Service".
398        collector.record_call("my.pkg.Greeter", "SayHello");
399
400        let report = collector.generate_json_report();
401        assert_eq!(report.summary.total, 1, "one method total");
402        assert_eq!(
403            report.summary.covered, 1,
404            "packaged service call should be counted as covered"
405        );
406
407        let text = collector.generate_text_report();
408        assert!(text.contains("100.0%"), "text report: {text}");
409    }
410
411    /// Build a pool with a message that nests two levels of sub-messages:
412    /// `Outer { id, inner: Inner }`, `Inner { name, addr: Addr }`,
413    /// `Addr { city }`. Recursively that is 5 fields (id, inner, inner.name,
414    /// inner.addr, inner.addr.city).
415    fn pool_with_nested_message() -> DescriptorPool {
416        use prost_reflect::prost_types::FieldDescriptorProto;
417        use prost_reflect::prost_types::field_descriptor_proto::{Label, Type};
418
419        let field =
420            |name: &str, number: i32, ty: Type, type_name: Option<&str>| FieldDescriptorProto {
421                name: Some(name.to_string()),
422                number: Some(number),
423                label: Some(Label::Optional as i32),
424                r#type: Some(ty as i32),
425                type_name: type_name.map(|s| s.to_string()),
426                ..Default::default()
427            };
428
429        let mut pool = DescriptorPool::new();
430        let file = FileDescriptorProto {
431            name: Some("nested.proto".to_string()),
432            message_type: vec![
433                DescriptorProto {
434                    name: Some("Outer".to_string()),
435                    field: vec![
436                        field("id", 1, Type::String, None),
437                        field("inner", 2, Type::Message, Some(".Inner")),
438                    ],
439                    ..Default::default()
440                },
441                DescriptorProto {
442                    name: Some("Inner".to_string()),
443                    field: vec![
444                        field("name", 1, Type::String, None),
445                        field("addr", 2, Type::Message, Some(".Addr")),
446                    ],
447                    ..Default::default()
448                },
449                DescriptorProto {
450                    name: Some("Addr".to_string()),
451                    field: vec![field("city", 1, Type::String, None)],
452                    ..Default::default()
453                },
454            ],
455            ..Default::default()
456        };
457        pool.add_file_descriptor_proto(file).unwrap();
458        pool
459    }
460
461    // Bug 2: the field-count denominator must recurse into nested messages so
462    // it matches the nested dotted paths counted as covered.
463    #[test]
464    fn nested_message_field_count_is_recursive() {
465        let pool = pool_with_nested_message();
466        let outer = pool.get_message_by_name("Outer").unwrap();
467        // Before the fix this returned 2 (only id, inner).
468        assert_eq!(CoverageCollector::count_fields_recursive(&outer), 5);
469    }
470
471    #[test]
472    fn nested_message_full_coverage_is_100_percent() {
473        let collector = CoverageCollector::new();
474        collector.register_pool(&pool_with_nested_message());
475        collector.record_fields_from_json(
476            "Outer",
477            &serde_json::json!({
478                "id": "x",
479                "inner": { "name": "y", "addr": { "city": "z" } }
480            }),
481        );
482
483        let report = collector.generate_json_report();
484        assert_eq!(report.field_summary.total, 5, "recursive field total");
485        assert_eq!(report.field_summary.covered, 5, "all nested fields covered");
486        let msg = report
487            .messages
488            .iter()
489            .find(|m| m.message_type == "Outer")
490            .unwrap();
491        assert_eq!(msg.total_fields, 5);
492        assert_eq!(msg.covered_fields.len(), 5);
493    }
494
495    #[test]
496    fn nested_message_partial_coverage_uses_recursive_total() {
497        let collector = CoverageCollector::new();
498        collector.register_pool(&pool_with_nested_message());
499        // Only 3 of the 5 recursive paths are exercised (id, inner, inner.name).
500        collector.record_fields_from_json(
501            "Outer",
502            &serde_json::json!({ "id": "x", "inner": { "name": "y" } }),
503        );
504
505        let report = collector.generate_json_report();
506        assert_eq!(report.field_summary.total, 5);
507        assert_eq!(report.field_summary.covered, 3);
508    }
509}