1use prost_reflect::{DescriptorPool, MessageDescriptor};
7use serde::{Deserialize, Serialize};
8use std::collections::{HashMap, HashSet};
9use std::sync::{Arc, Mutex};
10
11#[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#[derive(Debug, Clone, Serialize, Deserialize)]
26pub struct CoverageStats {
27 pub covered: usize,
28 pub total: usize,
29}
30
31#[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#[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#[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 fn count_fields_recursive(msg: &MessageDescriptor) -> usize {
123 fn count(msg: &MessageDescriptor, visited: &mut HashSet<String>) -> usize {
124 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 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 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 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 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 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 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 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 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 #[test]
394 fn coverage_matches_fully_qualified_service_name() {
395 let collector = CoverageCollector::new();
396 collector.register_pool(&pool_with_packaged_service());
397 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 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 #[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 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 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}